Revisiting Neural Scaling Laws in Language and Vision
Ibrahim Alabdulmohsin, Behnam Neyshabur, Xiaohua Zhai
Introduction
Scale has led to innovative research in both the vision domain and the natural language processing (NLP) domain. Recent work has found that scaling up the data size , the model size , the training schedule or all of them together often lead to improved performance. More importantly, scaling up the data size and the model size together can better utilize the compute resources. Scaling laws have been properly studied in several works, e.g. , and it has been found that the performance (e.g. excess loss) often follows a power law for some and as one varies a dimension of interest , such as the data or the model size.
While theoretical arguments alone seldom predict scaling law parameters in modern neural architectures , it has been observed that the benefit of scale could be predicted empirically . The general approach is to acquire a learning curve, i.e. a collection of samples , where is a dimension of interest such as the training data size while is a measure of performance, such as the validation loss. After that, parameters are estimated, e.g. by computing the best-fitting values of and in the model . Given the estimated scaling law parameters, one can then extrapolate by predicting for large values of .
Such learning curve extrapolation has found many applications, of which four seem to be more prominent. First, it offers a tool for understanding deep neural networks; e.g. how the architecture and data distribution impact scaling behaviors . Second, it has been used for sample size planning, particularly in data-scarce domains such as medicine . Third, it can reduce the environmental footprint of experimentation by terminating experiments early and accelerating hyper-parameter search . Forth, it has been applied in neural architecture search (NAS) . In addition, learning curve prediction offers a different methodology for comparing performance; e.g. instead of comparing accuracy on a single dataset, one can also examine the full (hypothetical) scaling curves.
However, in order to achieve such benefits in practice, it is imperative that scaling laws extrapolate accurately instead of merely interpolating the learning curve. To our knowledge, a validation of this sort based on extrapolation is often lacking in the literature and previous works have generally reported the best-fitting (interpolating) parameters. We demonstrate why this can be misleading in Section 4, where we illustrate how the scaling exponent that extrapolates best can be quite different from the exponent that best fits the given (finite) learning curve. In addition, we propose an estimator for scaling laws denoted , which extrapolates more accurately than previous methods as shown in Figure 1. We validate the proposed estimator in several domains, including image classification, neural machine translation, language modeling, and other related tasks.
We argue in Section 4 for a more rigorous methodology to validate scaling law parameters based on extrapolation, instead of only reporting the best-fitting (interpolating) parameters.
We propose a recipe in Section 3 to estimate scaling laws reliably from learning curves. The new estimator is verified across several domains: image classification, neural machine translation (NMT), language modeling, and tasks from BIG-Bench evaluation benchmark .
We use the proposed recipe to study the impact of the neural architecture’s type and size on scaling exponents.
We release a benchmark dataset consisting of 90 tasks to accelerate research in scaling laws.
Related work
Power law scaling in deep neural architectures has been verified in a wide range of domains, including image classification , language modeling , NMT , and speech recognition . To explain this theoretically, at least for data scaling, several works have argued for a power law behavior under various contexts. For instance, in the universal learning setting under the realizable case with a 0-1 misclassification loss, power law scaling emerges with exponent if the hypothesis space has an infinite Littlestone tree but not an infinite VC-Littlestone tree . Another argument for the exponent can be made in the non-realizable setting if the chosen loss is sufficiently smooth and the model size is limited by deriving the variance of the empirical solution around its population limit . A more relevant setting for deep neural networks is to assume that the model size is effectively infinite and the loss is Lipschitz continuous (e.g. continuous loss in bounded domains). Under the latter assumptions, it has been argued that the scaling exponent would satisfy where is the intrinsic dimension of the data manifold . This is consistent with the fact that scaling exponents often satisfy and that the exponent seems to be independent of the neural network architecture size as long as the architecture is sufficiently large .
Writing for the dimension of interest (e.g. data size) and for the error/loss of the model as a function of , three function classes have been used in the literature to model the performance as a function of while capturing its expected power law behavior:
The simplest model assumes a power law throughout the domain of : . This has been used, for example, to estimate the required sample size in healthcare , neural machine translation (NMT) , and language models , among others .
To capture saturating performance for large (i.e. when the Bayes optimal risk is bounded away from zero), a parameter is added: . This is, perhaps, the most commonly used model in the literature; see for instance .
A different parameterization has been recently used in NMT : , where and . Variants of this approach were used previously in studying, for example, scaling laws in vision transformers , accelerating hyper-parameter optimization , and (more generally) in learning curve prediction .
In this work, we introduce a fourth estimator and verify experimentally that it outperforms the above methods in terms of its extrapolation capability in several domains, as summarized in Figure 1. We describe and discuss its rationale in Section 3.
The function class , in which it is assumed that , captures (by definition) what it means for the excess risk to follow a power law. Hence, a question naturally arises: do we need any other function classes to estimate the scaling law parameters , and ?
Let be the learning curve and write for the restriction of the learning curve to . To extrapolate from a learning curve, we train each of the four models , , , and on the learning curve after applying a cutoff to mitigate the effect of small data samples. Then, we plot the excess risk predicted by each model, where is the (ground-truth) Bayes risk. Since is known exactly, an accurate model that extrapolates well would produce a linear curve in each plot. As shown in Figure 2, is accurate only when the data resides entirely in the power law regime (rightmost figure), whereas works well in all cases.
Derivation.
The function class arises from several natural requirements. First, we would like our function class to be sigmoid-like so that it fails only gracefully when the data deviates from the expected power law behavior; e.g. to avoid failures like that of and in Figure 2(left). Second, we would like our function class to reduce to power law functions as . More precisely, we require that:
To reiterate, this is because power law behavior has been empirically verified in a wide range of domains (see Section 2). Third, we would like our function class to be expressive enough to contain all of the functions in , i.e. , so that using becomes equivalent to using when the observed learning curve resides entirely in the power law regime.
If we take the first requirement above on the shape of the function, a general approach to achieve this is to write the performance as a convex combination of the form:
for some function that satisfies and . Here, is the predicted limiting performance when while is the performance at the random-guessing level. To meet the second requirement, we set for some learnable parameters and . Rearranging the terms yields Finally, we introduce a learnable parameter to meet our final requirement (see the ablation in Appendix A.1):
With , our function class reduces to as required. By differentiating both sides of the equation above and noting that , we deduce that remains a monotone decreasing function of for all as expected. In addition, by rearranging terms and using both the binomial theorem and the Lagrange series inversion theorem, we have the following asymptotic expansion for the excess loss as :
When , we recover the power law estimator . Setting allows to handle measurements that deviate from the power law behavior; i.e. when the learning curve does not fall into the asymptotic power law regime. The difference is (suppressing other constants).
The parameters to be fitted here are , , and . The parameter corresponds to the value of the loss at the random-guessing level and can be either fixed or optimized. We fix in our evaluation to be equal to the loss at the random-guessing level, although we observe similar results when it is optimized.
Loss Function.
In this work, scaling law parameters in all the four function classes are estimated by minimizing the square-log loss, similar to the approach used in . This serves two purposes. First, it penalizes the relative loss and, hence, treats errors at all scales equally. Second, it allows us to compute a subset of the parameters in closed-form using least squares.
Specifically, in , for example, we minimize:
In , we optimize the same loss above while fixing . In , we fix both and to zero. In , we optimize the following loss:
In all function classes, we use block coordinate descent, where we compute and in closed form using least squares, and estimate the remaining parameters (if any) using gradient descent, with a learning rate of . We repeat this until convergence.
Validating Scaling Laws using the Extrapolation Error
A common approach in the literature for estimating scaling law parameters is to assume a parametric model, e.g. , and reporting its best-fitting parameters to an empirical learning curve (see for example the prior works discussed in Section 2). Afterwards, patterns are reported about the behavior of the scaling law parameters; e.g. how the exponent varies with the architecture size. We argue, next, for a more rigorous methodology based on the extrapolation loss, instead of only reporting the best-fitting (interpolating) parameters. Specifically, choices of scaling law parameters that achieve a small interpolation error do not necessarily achieve a small extrapolation error so they may not, in fact, be valid estimates of scaling law parameters. Scaling law parameters should be validated by measuring how well they extrapolate.
To see why a validation of this sort matters, consider the following example. If we pretrain a vision transformer ViT/B/16 on subsets of JFT-300M (a proprietary dataset with 300M examples and 18k classes ) using the Adam optimizer with a base learning rate of 5e-4, batch-size 4,096, and dropout rate of 0.1, and evaluate the 10-shot error rate on ImageNet-ILSRCV2012 , we obtain the learning curve shown in Figure 3(left, in green). Evidently, power law emerges; i.e. ImageNet 10-shot error rate (shown in green) follows a linear curve on a log-log plot. Hence, one might estimate, for example the scaling exponent using least squares.
However, consider now the family of curves shown in Figure 3(left), all corresponding to but with scaling exponents that vary from about to (while fitting the parameters and ). Note that all five curves overlap with each other significantly. Choosing the best fitting parameters on the learning curve would favor a scaling exponent of as shown in Figure 3(right). However, if we validate the parameters by evaluating how well they extrapolate (i.e. how well they predict performance when the number of seen examples ), a different picture emerges. We observe that a more accurate estimate of the scaling exponent is . This is shown in Figure 3(left) and in Figure 3(right) by measuring the extrapolation loss. Here, validation is measured using the root mean square error (RMSE) to the log-loss:
in which x is uniform over the set , where is the predicted loss while is the actual. We apply the logarithm so that we penalize the relative error and, hence, assign equal importance to all error scales (both large and small)E.g. if we have a single measurement where and the estimator predicts , the RMSE in (6) is approximately equal to 1%. Similarly, it is approximately equal to 1% when while ..
In summary, scaling law parameters that give the best fit on the learning curve (i.e. lowest interpolation loss) do not generally extrapolate best. When using extrapolation loss instead, different scaling law parameters emerge. In this work, we use extrapolation to evaluate the quality of scaling law estimators.
Experiments
We provide an empirical evaluation of the four scaling law estimators in several domains, including image classification (72 tasks), neural machine translation (5 tasks), language modeling (5 tasks), and other language-related evaluations (10 tasks). The dataset for neural machine translation is available at . The code and dataset for the remaining tasks used in this evaluation are made publicly available to facilitate further research in this domain Code and benchmark dataset will be made available at: https://github.com/google-research/google-research/tree/master/revisiting_neural_scaling_laws.. In all experiments, we divide the learning curve into two splits: (1) one split used for training the scaling law estimators, and (2) one split used for evaluating extrapolation. Setting , where is the maximum value of in the data, the first split is the domain while the second split is the domain . We measure extrapolation error using RMSE in (6). All experiments are executed on Tensor Processing Units (TPUs).
We use three families of architectures: (1) big-transfer residual neural networks (BiT) , (2) vision transformers (ViT) , and (3) MLP mixers (MiX) . For each family, we have two models of different sizes as shown in Table 1 in order to assess the impact of the size of the architecture on the scaling parameters. We pretrain on JFT-300M . Since pre-training task performance is not representative of the downstream performance , we evaluate the few-shot accuracy downstream on four datasets: (1) ImageNet , (2) Birds 200 , (3) CIFAR100 , and (4) Caltech101 . For each dataset, we report 5/10/25-shot accuracy. It results in 72 tasks for the combinations of architecture, dataset, and metric. Following , we removed duplicate pre-training examples between upstream JFT-300M dataset and all the downstream train and test sets.
Bootstrapped Examples.
In the few-shot image classification setting under the transfer learning setup, overfitting can occur if the upstream dataset is small, where training beyond a particular number of steps would reduce the downstream validation accuracy . This is demonstrated in Figure 4, where we pretrain on subsets of JFT-300M (upstream) and evaluate ImageNet 10-shot error (downstream).
Nevertheless, we observe that prior to reaching peak performance, training examples behave as if they were fresh samples. This observation generalizes the bootstrapping phenomenon observed in , where it showed that training examples behave as fresh samples prior to convergence, which would be equivalent to our observation if no overfitting occurs. Throughout the sequel, we refer to the examples seen during training prior to peak performance as “bootstrapped examples" and use their number as the independent variable when evaluating the scaling law estimators in this section.
Figure 5 illustrates how well each of the the four scaling law estimators , , , and can extrapolate from a given learning curve. The complete set of figures is provided in Appendix A.3. We observe that extrapolates better than other methods and produces learning curves that approximate the empirical results more faithfully. As shown in Figure 1, outperforms the other methods in more than 70% of the tasks in this domain.
Impact of the Architecture.
Figure 6 plots the scaling exponent in each architecture when the downstream task is -shot accuracy on ImageNet. We observe that within each family of models, larger models have more favorable scaling exponents. In addition, yields estimates of the scaling exponents that are larger in absolute magnitude than in other methods. Figure 7 shows that such differences in scaling exponents show up indeed in the slopes of the learning curves as expected.
2 Neural Machine Translation (NMT)
Next, we evaluate the scaling law estimators on NMT. We use the setup studied in , in which models are trained with the per-token cross-entropy loss using Adafactor optimizer with a batch-size of 500K tokens and a dropout rate of 0.1 . We use the encoder-decoder transformer models 6L6L, 28L6L and 6L28L, where 28L6L means that the architecture has 28 encoders and 6 decoders. We also use the two architectures: decoder-only with language modeling loss (D/LM) and the transformer-encoder with LSTM decoder (TE/LSTM). These correspond to the architectures used in Figure 1 in . In all cases, performance is measured using log-perplexity on a hold-out dataset. To evaluate the accuracy of the scaling law estimator, we fit its parameters on the given learning curve (for up to 256M sentence pairs) and use it to predict the log-perplexity when the architecture is trained on 512M sentence pairs. Because the learning curves contain few points, we only evaluate on the 512M sentence pairs. Table 2 displays the RMSE of each estimator. Clearly, performs better than the other methods as summarized in Figure 1, which is consistent with the earlier results in image classification.
3 Language Modeling
Next, we evaluate the scaling law estimators in language modeling, where the goal is to predict the next token. We use the LaMDA architecture used in , which is a decoder-only transformer language model. Five model sizes are used, ranging from to model parameters. In each model, we rescale validation loss to the unit interval $\mathcal{M}_{4}\mathcal{M}_{2}\mathcal{M}_{4}\mathcal{M}_{4}\mathcal{M}_{2}c\mathcal{M}_{1}\mathcal{M}_{3}c\mathcal{M}_{2}\mathcal{M}_{4}-1/3$ and decreases (in absolute magnitude) for larger models.
4 Scalable Tasks from the BIG-Bench Evaluation Benchmark
Finally, we evaluate the scaling law estimators on language tasks from the BIG-Bench collaborative benchmark . Here, we pretrain a 262M-parameter decoder-only transformer (middle architecture in Figure 8) on language modeling and evaluate its 1-shot and 2-shot capabilities in five language-related tasks. We choose the five tasks that exhibit the highest learnability from the benchmark (i.e. improvement in performance when pretrained on language modeling, see for details). The five tasks are: linguistic_mappings, qa_wikidata, unit_conversion, mult_data_wrangling and date_understanding. We use the benchmark’s preferred metrics in all cases, which is either “multiple choice grade" or “exact string match" depending on the task. Table 3 and Figure 1 summarize the results. In this evaluation, both and perform best and equally well. In addition, we observe that and perform equally well and consistently worse than the other methods. One possible reason is that the learning curves are quite noisy (see Appendix A.2).
Discussion
The remarkable progress in deep learning in recent years is largely driven by improvements in scale, where bigger models are trained on larger datasets for longer training schedules. Several works observe that the benefit of scale can be predicted empirically by extrapolating from learning curves and this has found important applications, such as in sample size planning and neural architecture search. However, to achieve such benefits in practice, it is imperative that scaling laws extrapolate accurately. We demonstrate that scaling parameters that yield the best fit to the learning curve do not generally extrapolate best, thereby challenging their use as valid estimate of scaling law parameters. Hence, we argue for a more rigorous validation of scaling law parameters based on the extrapolation loss. In addition, we present a recipe for estimating scaling law parameters that extrapolates more accurately than in previous works, which we verify in several state-of-the-art architecture across a wide range of domains. To facilitate research in this domain, we also release a benchmark dataset comprising of 90 evaluation tasks. We believe that the proposed scaling law estimator can be utilized, for example, to accelerate neural architecture search (NAS), which we plan to study in future work.
Acknowledgements
The authors would like to acknowledge and thank Behrooz Ghorbani for his feedback on earlier drafts of this manuscript as well as Ambrose Slone and Lucas Beyer for their help with some of the experiments. We also would like to thank Daniel Keysers and Olivier Bousquet for the useful discussions.
References
Checklist
Do the main claims made in the abstract and introduction accurately reflect the paper’s contributions and scope? [Yes]
Did you describe the limitations of your work? [N/A] We are aware of many applications using learning curve extrapolation, and plan to apply our new recipe to new applications (e.g. NAS) to land the impact.
Did you discuss any potential negative societal impacts of your work? [N/A] We provide a study of scaling laws and an improved estimator of scaling law parameters. We do not anticipate any negative societal impacts of this work.
Have you read the ethics review guidelines and ensured that your paper conforms to them? [Yes]
If you are including theoretical results…
Did you state the full set of assumptions of all theoretical results? [N/A]
Did you include complete proofs of all theoretical results? [N/A]
Did you include the code, data, and instructions needed to reproduce the main experimental results (either in the supplemental material or as a URL)? [Yes] We plan to release the code and a benchmark dataset to facilitate research in this direction. Some of the datasets used in our experiments are proprietary and cannot be released, such as JFT-300M. We also include experiments on publicly available datasets, such as Big-Bench for reproducibility.
Did you specify all the training details (e.g., data splits, hyperparameters, how they were chosen)? [Yes] We describe the data, training procedure and architectures in details. See Section 5.
Did you report error bars (e.g., with respect to the random seed after running experiments multiple times)? [Yes] Please see Tables 3 and 2.
Did you include the total amount of compute and the type of resources used (e.g., type of GPUs, internal cluster, or cloud provider)? [Yes] All experiments are executed in Tensor Processing Units (TPUs). See Section 5.
If you are using existing assets (e.g., code, data, models) or curating/releasing new assets…
If your work uses existing assets, did you cite the creators? [Yes] We cite all the datasets and architectures we use.
Did you mention the license of the assets? [No]
Did you include any new assets either in the supplemental material or as a URL? [Yes] We plan to release a code and a benchmark dataset.
Did you discuss whether and how consent was obtained from people whose data you’re using/curating? [N/A]
Did you discuss whether the data you are using/curating contains personally identifiable information or offensive content? [N/A]
If you used crowdsourcing or conducted research with human subjects…
Did you include the full text of instructions given to participants and screenshots, if applicable? [N/A]
Did you describe any potential participant risks, with links to Institutional Review Board (IRB) approvals, if applicable? [N/A]
Did you include the estimated hourly wage paid to participants and the total amount spent on participant compensation? [N/A]
Appendix A Appendix
In this section, we show that the improvement in compared to is due to both of the conditions discussed in Section 3: (1) has a sigmoid-like shape so that it can handle deviations from the power law behavior and (2) it contains all of the functions in so that reduces to when the learning curve falls entirely in the power law regime.
First, we observe that the second condition alone is not sufficient to have a high extrapolation accuracy since itself has a less extrapolation accuracy than . To show that a sigmoid-like behavior is not sufficient, we evaluation a different version of , in which is not introduced. Recall that the parameter was introduced so that contains . Figure 10 and Table 4 present the extrapolation performance of all the estimators when using 10-shot ImageNet accuracy as a metric. We observe that without the parameter does not perform as well as when is introduced.
A.2 Big Bench Learning Curves
In Figure 11, we plot the learning curves in the BIG-bench evaluation tasks. We observe that the learning curves are noisier in this setting than in previous cases.