Beyond neural scaling laws: beating power law scaling via data pruning

Ben Sorscher, Robert Geirhos, Shashank Shekhar, Surya Ganguli, Ari S. Morcos

Introduction

Empirically observed neural scaling laws in many domains of machine learning, including vision, language, and speech, demonstrate that test error often falls off as a power law with either the amount of training data, model size, or compute. Such power law scaling has motivated significant societal investments in data collection, compute, and associated energy consumption. However, power law scaling is extremely weak and unsustainable. For example, a drop in error from 3%3\% to 2%2\% might require an order of magnitude more data, compute, or energy. In language modeling with large transformers, a drop in cross entropy loss from about 3.4 to 2.8 natsHowever, note that nats is on a logarithmic scale and and small improvements in nats can lead to large improvements in downstream tasks. requires 10 times more training data (Fig. 1 in ). Also, for large vision transformers, an additional 22 billion pre-training data points (starting from 11 billion) leads to an accuracy gain on ImageNet of a few percentage points (Fig. 1 in ). Here we ask whether we might be able to do better. For example, can we achieve exponential scaling instead, with a good strategy for selecting training examples? Such vastly superior scaling would mean that we could go from 3%3\% to 2%2\% error by only adding a few carefully chosen training examples, rather than collecting 10×10\times more random ones.

Focusing on scaling of performance with training dataset size, we demonstrate that exponential scaling is possible, both in theory and practice. The key idea is that power law scaling of error with respect to data suggests that many training examples are highly redundant. Thus one should in principle be able to prune training datasets to much smaller sizes and train on the smaller pruned datasets without sacrificing performance. Indeed some recent works have demonstrated this possibility by suggesting various metrics to sort training examples in order of their difficulty or importance, ranging from easy or redundant examples to hard or important ones, and pruning datasets by retaining some fraction of the hardest examples. However, these works leave open fundamental theoretical and empirical questions: When and why is successful data pruning possible? What are good metrics and strategies for data pruning? Can such strategies beat power law scaling? Can they scale to ImageNet? Can we leverage large unlabeled datasets to successfully prune labeled datasets? We address these questions through both theory and experiment. Our main contributions are:

Employing statistical mechanics, we develop a new analytic theory of data pruning in the student-teacher setting for perceptron learning, where examples are pruned based on their teacher margin, with large (small) margins corresponding to easy (hard) examples. Our theory quantitatively matches numerical experiments and reveals two striking predictions:

The optimal pruning strategy changes depending on the amount of initial data; with abundant (scarce) initial data, one should retain only hard (easy) examples.

Exponential scaling is possible with respect to pruned dataset size provided one chooses an increasing Pareto optimal pruning fraction as a function of initial dataset size.

We show that the two striking predictions derived from theory hold also in practice in much more general settings. Indeed we empirically demonstrate signatures of exponential scaling of error with respect to pruned dataset size for ResNets trained from scratch on SVHN, CIFAR-10 and ImageNet, and Vision Transformers fine-tuned on CIFAR-10.

Motivated by the importance of finding good quality metrics for data pruning, we perform a large scale benchmarking study of 10 different data pruning metrics at scale on ImageNet, finding that most perform poorly, with the exception of the most compute intensive metrics.

We leveraged self-supervised learning (SSL) to developed a new, cheap unsupervised data pruning metric that does not require labels, unlike prior metrics. We show this unsupervised metric performs comparably to the best supervised pruning metrics that require labels and much more compute. This result opens the door to the exciting possibility of leveraging pre-trained foundation models to prune new datasets even before they are labeled.

Overall these results shed theoretical and empirical insights into the nature of data in deep learning and our ability to prune it, and suggest our current practice of collecting extremely large datasets may be highly inefficient. Our initial results in beating power law scaling motivate further studies and investments in not just inefficently collecting large amounts of random data, but rather, intelligently collecting much smaller amounts of carefully selected data, potentially leading to the creation and dissemination of foundation datasets, in addition to foundation models .

Background and related work

Our work brings together 33 largely disparate strands of intellectual inquiry in machine learning: (1) explorations of different metrics for quantifying differences between individual training examples; (2) the empirical observation of neural scaling laws; and (3) the statistical mechanics of learning.

Several recent works have explored various metrics for quantifying individual differences between data points. To describe these metrics in a uniform manner, we will think of all of them as ordering data points by their difficulty, ranging from “easiest” to “hardest.” When these metrics have been used for data pruning, the hardest examples are retained, while the easiest ones are pruned away.

For example trained small ensembles (of about 1010) networks for a very short time (about 1010 epochs) and computed for every training example the average L2L_{2} norm of the error vector (EL2N score). Data pruning by retaining only the hardest examples with largest error enabled training from scratch on only 50%50\% and 75%75\% of CIFAR-10 and CIFAR-100 respectively without any loss in final test accuracy. However the performance of EL2N on ImageNet has not yet been explored.

noticed that over the entire course of training, some examples are learned early and never forgotten, while others can be learned and unlearned (i.e. forgotten) repeatedly. They developed a forgetting score which measures the degree of forgetting of each example. Intuitively examples with low (high) forgetting scores can be thought of as easy (hard) examples. explored data pruning using these metrics, but not at ImageNet scale.

defined a memorization score for each example, corresponding to how much the probability of predicting the correct label for the example increases when it is present in the training set relative to when it is absent (also see ); a large increase means the example must be memorized (i.e. the remaining training data do not suffice to correctly learn this example). Additionally also considered an influence score that quantifies how much adding a particular example to the training set increases the probability of the correct class label of a test example. Intuitively, low memorization and influence scores correspond to easy examples that are redundant with the rest of the data, while high scores correspond to hard examples that must be individually learned. did not use these scores for data pruning as their computation is expensive. We note since memorization explicitly approximates the increase in test loss due to removing each individual example, it is likely to be a good pruning metric (though it does not consider interactions).

Active learning iterates between training a model and selecting new inputs to be labeled . In contrast, we focus on data pruning: one-shot selection of a data subset sufficient to train to high accuracy from scratch. A variety of coreset algorithms (e.g. ) have been proposed for this, but their computation is expensive, and so data-pruning has been less explored at scale on ImageNet. An early clustering approach allowed training on 90%90\% of ImageNet without sacrificing accuracy. Notably reduced this to 80%80\% by training a large ensemble of networks on ImageNet and using ensemble uncertainty to define the difficulty of each example, with low (high) uncertainty corresponding to easy (hard) examples. We will show how to achieve similar pruning performance without labels or the need to train a large ensemble.

assigned a score to every ImageNet image, given by the number of models in a diverse ensemble (10 models) that misclassified the image. Intuitively, low (high) scores correspond to easy (hard) examples. The pruning performance of this metric remains unexplored.

We note: (1) only one of these metrics has tested well for its efficacy in data pruning at scale on ImageNet; (2) all of these metrics require label information; (3) there is no theory of when and why data pruning is possible for any of these metrics; and (4) none of these works suggest the possibility of exponential scaling. We thus go beyond this prior work by benchmarking the data pruning efficacy of not only these metrics but also a new unsupervised metric we introduce that does not require label information, all at scale on ImageNet. We also develop an analytic theory for data-pruning for the margin metric that predicts not only the possibility of exponential scaling but also the novel finding that retaining easy instead of hard examples is better when data is scarce.

2 Neural scaling laws and their potential inefficiency

Recent work has demonstrated that test loss L\mathcal{L} often falls off as a power law with different resources like model parameters (NN), number of training examples (PP), and amount of compute (CC). However, the exponents ν\nu of these power laws are often close to , suggesting potentially inefficient use of resources. For example, for large models with lots of compute, so that the amount of training data constitutes a performance bottleneck, the loss scales as L≈P−ν\mathcal{L}\approx P^{-\nu}. Specifically for a large transformer based language model, ν=0.095\nu=0.095, which implies an order of magnitude increase in training data drops cross-entropy loss by only about 0.60.6 nats (Fig. 1 in ). In neural machine translation experiments ν\nu varies across language pairs from 0.350.35 to 0.480.48 (Table 1 in ). Interestingly, explored a fixed computation budget CC and optimized jointly over model size NN and training set size PP, revealing that scaling both NN and PP commensurately as CC increases is compute optimal, and can yield smaller high performing models (trained on more data) than previous work. Nevertheless, for a transformer based language model, a 100×100\times increase in compute, corresponding to 10×10\times increases in both model size and training set size, leads to a drop in cross-entropy loss of only about 0.50.5 nats (Fig. 2 in ). Similar slow scaling holds for large vision transformers where adding 22 billion pre-training images reduces ImageNet performance by a few percentage points (Fig. 1 in ). While all of these results constitute significant improvements in performance, they do come at a substantial resource cost whose fundamental origin arises from power law scaling with small exponents. Recent theoretical works have argued that the power law exponent is governed by the dimension of a data manifold from which training examples are uniformly drawn. Here we explore whether we can beat power law scaling through careful data selection.

3 Statistical mechanics of perceptron learning

Statistical mechanics has long played a role in analyzing machine learning problems (see e.g. for reviews). One of the most fundamental applications is perceptron learning in the student-teacher setting , in which random i.i.d. Gaussian inputs are labeled by a teacher perceptron to construct a training set. The test error for another student perceptron learning from this training set then scales as a power law with exponent −1-1 for such data. Such perceptrons have also been analyzed in an active learning setting where the learner is free to design any new input to be labeled , rather than choose from a fixed set of inputs, as in data-pruning. Recent work has analyzed this scenario but focused on message passing algorithms that are tailored to the case of Gaussian inputs and perceptrons, and are hard to generalize to real world settings. In contrast we analyze margin based pruning algorithms that are used in practice in diverse settings, as in .

An analytic theory of data pruning

We are interested in the test error ε\varepsilon of this final perceptron as a function of αtot\alpha_{\text{tot}}, ff, and the angle θ\theta between the probe student \mbox{\mathbf{J}}_{\text{probe}} and the teacher T\mathbf{T}. Our theory approximates \mbox{\mathbf{J}}_{\text{probe}} as simply a random Gaussian vector conditioned to have angle θ\theta with the teacher T\mathbf{T}. Under this approximation we obtain an analytic theory for ε(αtot,f,θ)\varepsilon(\alpha_{\text{tot}},f,\theta) that is asymptotically exact in the high dimensional limit (App. A). We first examine results when θ=0\theta=0, so we are pruning training examples according to their veridical margins with respect to the teacher (Fig. 1A). We find two striking phenomena, each of which constitute predictions in real-world settings that we will successfully confirm empirically.

First, we note the test error curve for f=1f=1 in Fig. 1A corresponding to no pruning, or equivalently to randomly pruning a larger dataset of size αtot\alpha_{\text{tot}} down to a size αprune\alpha_{\text{prune}}, exhibits the well known classical perceptron learning power law scaling ε∝αprune−1\varepsilon\propto\alpha_{\text{prune}}^{-1}. Interestingly though, for small αtot\alpha_{\text{tot}}, keeping the hardest examples performs worse than random pruning (lighter curves above darkest curve for small αprune\alpha_{\text{prune}} in Fig. 1A). However, for large αtot\alpha_{\text{tot}}, keeping the hardest examples performs substantially better than random pruning (lighter curves below darkest curve for large αprune\alpha_{\text{prune}} in Fig. 1A). It turns out keeping the easiest rather than hardest examples is a better pruning strategy when αtot\alpha_{\text{tot}} is small (Fig. 1C). If one does not have much data to start with, it is better to keep the easiest examples with largest margins (i.e. the blue regions of Fig. 1B) to avoid overfitting. The easiest examples provide coarse-grained information about the target function, while the hard examples provide fine-grained information about the target function which can prevent the model from learning if one starts with lots of data. In cases where overfitting is less of an issue, it is best to keep the hardest examples with smallest margin that provide more information about the teacher’s decision boundary (i.e. the green region of Fig. 1B). Intuitively, in the limited data regime, it is challenging to model outliers since the basics are not adequately captured; hence, it is more important to keep easy examples so that the model can get to moderate error. However, with a larger dataset, the easy examples can be learned without difficulty, making modeling outliers the fundamental challenge.

Fig. 1C reveals which pruning strategy is best as a joint function of αtot\alpha_{\text{tot}} and ff. Note the transition between optimal strategies becomes sharper at small fractions ff of data kept. This transition between optimal pruning strategies can be viewed as a prediction in more general settings. To test this prediction we trained a ResNet18 on pruned subsets of the CIFAR-10 dataset (Fig. 1D), and observed strikingly similar behavior, indicating the prediction can hold far more generally, beyond perceptron learning. Interestingly, missed this transition, likely because they started pruning from large datasets.

A second prediction of our theory is that when keeping a fixed fraction ff of the hardest examples as αtot\alpha_{\text{tot}} increases (i.e. constant color curves in Fig. 1A), the error initially drops exponentially in αprune=fαtot\alpha_{\text{prune}}=f\alpha_{\text{tot}}, but then settles into the universal power law ε∝αprune−1\varepsilon\propto\alpha_{\text{prune}}^{-1} for all fixed ff. Thus there is no asymptotic advantage to data pruning at a fixed ff. However, by pruning more aggressively (smaller ff) when given more initial data (larger αtot\alpha_{\text{tot}}), one can achieve a Pareto optimal test error as a function of pruned dataset size αprune\alpha_{\text{prune}} that remarkably traces out at least an exponential scaling law (Fig. 1A, purple curve). Indeed our theory predicts for each αprune\alpha_{\text{prune}} a Pareto optimal point in αtot\alpha_{\text{tot}} and ff (subject to αprune=fαtot\alpha_{\text{prune}}=f\alpha_{\text{tot}}), yielding for every fixed αprune\alpha_{\text{prune}} an optimal foptf_{\text{opt}}, plotted in Fig. 1E. Note foptf_{\text{opt}} decreases with αprune\alpha_{\text{prune}} indicating more aggressive pruning (smaller foptf_{\text{opt}}) of original datasets of larger size αtot\alpha_{\text{tot}} is required to obtain larger Pareto optimal pruned datasets of size αprune\alpha_{\text{prune}}. We will test this striking scaling prediction in Fig. 3.

Classical randomly selected data generates slow power law error scaling because each extra training example provides less new information about the correct decision boundary than the previous example. More formally, let S(αtot)S(\alpha_{\text{tot}}) denote the typical entropy of the posterior distribution over student perceptron weights consistent with a training set of size αtot\alpha_{\text{tot}}. The information gain I(αtot)I(\alpha_{\text{tot}}) due to additional examples beyond αtot\alpha_{\text{tot}} can be defined as the rate at which the posterior entropy is reduced: I(αtot)=−ddαtotS(αtot)I(\alpha_{\text{tot}})=-\frac{d}{d\alpha_{\text{tot}}}S(\alpha_{\text{tot}}). In classical perceptron learning I(αtot)I(\alpha_{\text{tot}}) decays to zero as a power law in αtot\alpha_{\text{tot}}, reflecting a vanishing amount of information per each new example, leading to the slow power law decay of test error ε∝αtot−1\varepsilon\propto\alpha_{\text{tot}}^{-1}. However, data pruning can increase the information gained per example by pruning away the uninformative examples. To show this, we generalized the replica calculation of the posterior entropy SS and information gain II from random datasets of size αtot\alpha_{\text{tot}} to pruned datasets of size αprune\alpha_{\text{prune}} (App. A). We plot the resulting information gain I(αprune)I(\alpha_{\text{prune}}) for different ff in Fig. 1F. For any fixed ff, I(αprune)I(\alpha_{\text{prune}}) will eventually decay as a power law as αprune−1\alpha_{\text{prune}}^{-1}. However, by more aggressively pruning (smaller ff) datasets of larger size αtot\alpha_{\text{tot}}, I(αprune)I(\alpha_{\text{prune}}) can converge to a finite value I(∞)=1I(\infty)=1 nat/example, resulting in larger pruned datasets only adding useful non-redundant information. Since each new example under Pareto optimal data pruning conveys finite information about the target decision boundary, as seen in Fig. 1F, the test error can decay at least exponentially in pruned dataset size as in Fig. 1A. Classical results have shown that training examples chosen by maximizing the disagreement of a committee of student perceptrons can provide an asymptotically finite information rate, leading to exponential decay in test error. Intriguingly, the Pareto-optimal data pruning strategy we study in this work leads to faster than exponential decay, because it includes (partial) information about the target function provided by the probe student (Fig. 11).

We next examine the case of nonzero angle θ\theta between the probe student \mbox{\mathbf{J}}_{\text{probe}} and the teacher T\mathbf{T}, such that the ranking of training examples by margin is no longer completely accurate (Fig. 2A). Retaining the hard examples with smallest margin with respect to the probe student will always result in pruned datasets lying near the probe’s decision boundary. But if θ\theta is large, such examples might be far from the teacher’s decision boundary, and therefore could be less informative about the teacher (Fig. 2A). As a result our theory, confirmed by simulations, predicts that under nonzero angles θ\theta, the Pareto optimal lower envelope of test error over both αtot\alpha_{\text{tot}} and ff initially scales exponentially as a function of αprune=fαtot\alpha_{\text{prune}}=f\alpha_{\text{tot}} but then crosses over to a power law (Fig. 2BCD). Indeed, at any given nonzero θ\theta, our theory reveals that as αtot\alpha_{\text{tot}} (and therefore αprune\alpha_{\text{prune}}) becomes large, one cannot decrease test error any further by retaining less than a minimum fraction fmin(θ)f_{\text{min}}(\theta) of all available data. For example when θ=10∘\theta=10^{\circ} (θ=20∘\theta=20^{\circ}) one can do no better asymptotically than pruning down to 24% (46%46\%) of the total data (Fig. 2CD). As θ\theta approaches , fmin(θ)f_{\text{min}}(\theta) approaches , indicating that one can prune extremely aggressively to arbitrarily small ff while still improving performance, leading to at least exponential scaling for arbitrarily large αprune\alpha_{\text{prune}} in Fig. 2B. However, for nonzero θ\theta, the lack of improvement for f<fmin(θ)f<f_{\text{min}}(\theta) at large αprune\alpha_{\text{prune}} renders aggressive pruning ineffective. This result highlights the importance of finding high quality pruning metrics with θ≈0\theta\approx 0. Such metrics can delay the cross over from exponential to power law scaling as pruned dataset size αprune\alpha_{\text{prune}} increases, by making aggressive pruning with very small ff highly effective. Strikingly, in App. Fig. 10 we demonstrate this cross-over in a real-world setting by showing that the test error on SVHN is bounded below by a power law when the dataset is pruned by a probe ResNet18 under the EL2N metric, trained for 4 epochs (weak pruning metric) but not a probe ResNet18 trained for 40 epochs (strong pruning metric).

Data pruning can beat power law scaling in practice

Our theory of data pruning for the perceptron makes three striking predictions which can be tested in more general settings, such as deep neural networks trained on benchmark datasets: (1) relative to random data pruning, keeping only the hardest examples should help when the initial dataset size is large, but hurt when it is small; (2) data pruning by retaining a fixed fraction ff of the hardest examples should yield power law scaling, with exponent equal to that of random pruning, as the initial dataset size increases; (3) the test error optimized over both initial data set size and fraction of data kept can trace out a Pareto optimal lower envelope that beats power law scaling of test error as a function of pruned dataset size, through more aggressive pruning at larger initial dataset size. We verified all three of these predictions on ResNets trained on SVHN, CIFAR-10, and ImageNet using varying amounts of initial dataset size and fractions of data kept under data pruning (compare theory in Fig. 3A with deep learning experiments in Fig. 3BCD). In each experimental setting we see better than power law scaling at larger initial data set sizes and more aggressive pruning. Moreover we would likely see even better scaling with even larger initial datasets (as in Fig.3A dashed lines).

Modern foundation models are pre-trained on a large initial dataset, and then transferred to other downstream tasks by fine-tuning on them. We therefore examined whether data-pruning can be effective for both reducing the amount of fine-tuning data and the amount of pre-training data. To this end, we first analyzed a vision transformer (ViT) pre-trained on ImageNet21K and then fine-tuned on different pruned subsets of CIFAR-10. Interestingly, pre-trained models allow for far more aggressive data pruning; fine-tuning on only 10% of CIFAR-10 can match or exceed performance obtained by fine tuning on all of CIFAR-10 (Fig. 4A). Furthermore Fig. 4A provides a new example of beating power law scaling in the setting of fine-tuning. Additionally, we examined the efficacy of pruning pre-training data by pre-training ResNet50s on different pruned subsets of ImageNet1K (exactly as in Fig. 3D) and then fine-tuning them on all of CIFAR-10. Fig. 4B demonstrates pre-training on as little as 50%50\% of ImageNet can match or exceed CIFAR-10 performance obtained by pre-training on all of ImageNet. Thus intriguingly pruning pre-training data on an upstream task can still maintain high performance on a different downstream task. Overall these results demonstrate the promise of data pruning in transfer learning for both the pre-training and fine-tuning phases.

Benchmarking supervised pruning metrics on ImageNet

We note that the majority of data pruning experiments have been performed on small-scale datasets (i.e. variants of MNIST and CIFAR), while the few pruning metrics proposed for ImageNet have rarely been compared against baselines designed on smaller datasets. Therefore, it is currently unclear how most pruning methods scale to ImageNet and which method is best. Motivated by how strongly the quality of a pruning metric can impact performance in theory (Fig. 2), we decided to fill this knowledge gap by performing a systematic evaluation of 88 different supervised pruning metrics on ImageNet: two variants of influence scores , two variants of EL2N , DDD , memorization , ensemble active learning , and forgetting . See Section 2 for a review of these metrics. Additionally, we include two new prototypicality metrics that we introduce in the next section.

We first asked how consistent the rankings induced by different metrics are by computing the Spearman rank correlation between each pair of metrics (Fig. 5A). Interestingly, we found substantial diversity across metrics, though some (EL2N, DDD, and memorization) were fairly similar with rank correlations above 0.70.7. However, we observed marked performance differences between metrics: Fig 5BC shows test performance when a fraction ff of the hardest examples under each metric are kept in the training set. Despite the success of many of these metrics on smaller datasets, only a few still match performance obtained by training on the full dataset, when selecting a significantly smaller training subset (i.e. about 80%80\% of ImageNet). Nonetheless, most metrics continue to beat random pruning, with memorization in particular demonstrating strong performance (Fig. 5C). We note that data pruning on ImageNet may be more difficult than data pruning on other datasets, because ImageNet is already carefully curated to filter out uninformative examples.

We found that all pruning metrics amplify class imbalance, which results in degraded performance. To solve this we used a simple 50%50\% class balancing ratio for all ImageNet experiments. Further details and baselines without class balancing are shown in App. H. Metric scores, including baselines, are available from https://github.com/rgeirhos/dataset-pruning-metrics.

Self-supervised data pruning through a prototypicality metric

Fig. 5 shows many data pruning metrics do not scale well to ImageNet, while the few that do require substantial amounts of compute. Furthermore, all these metrics require labels, thereby limiting their ability to prune data for large-scale foundation models trained on massive unlabeled datasets . Thus there is a clear need for simple, scalable, self-supervised pruning metrics.

To compute a self-supervised pruning metric for ImageNet, we perform kk-means clustering in the embedding space of an ImageNet pre-trained self-supervised model (here: SWaV ), and define the difficulty of each data point by the cosine distance to its nearest cluster centroid, or prototype. Thus easy (hard) examples are the most (least) prototypical. Encouragingly, in Fig. 5C, we find our self-supervised prototype metric matches or exceeds the performance of the best supervised metric, memorization, until only 70–80% of the data is kept, despite the fact that our metric does not use labels and is much simpler and cheaper to compute than many previously proposed supervised metrics. See App. Fig. 9 for further scaling experiments using the self-supervised metric.

To assess whether the clusters found by our metric align with ImageNet classes, we compared their overlaps in Fig. 6A. Interestingly, we found alignment for some but not all classes. For example, class categories such as snakes were largely aligned to a small number of unsupervised clusters, while other classes were dispersed across many such clusters. If class information is available, we can enforce alignment between clusters and classes by simply computing a single prototype for each class (by averaging the embeddings of all examples of this class). While originally intended to be an additional baseline metric (called supervised prototypes, light blue in Fig 5C), this metric remarkably outperforms other supervised metrics and largely matches the performance of memorization, which is prohibitively expensive to compute. Moreover, the performance of the best self-supervised and supervised metrics are similar, demonstrating the promise of self-supervised pruning.

One important choice for the self-supervised prototype metric is the number of clusters kk. We found, reassuringly, our results were robust to this choice: kk can deviate one order of magnitude more or less than the true number of classes (i.e. 10001000 for ImageNet) without affecting performance (App. F).

To better understand example difficulty under various metrics, we visualize extremal images for our self-supervised prototype metric and the memorization metric for one class (Fig 6B,C). Qualitatively, easy examples correspond to highly similar, redundant images, while hard examples look like idiosyncratic outliers. See App. E, Figs. 12,13,14,15,16,17,18,19 for more classes and metrics.

Discussion

We have shown, both in theory and practice, how to break beyond slow power law scaling of error versus dataset size to faster exponential scaling, through data pruning. Additionally we have developed a simple self-supervised pruning metric that enables us to discard 20% of ImageNet without sacrificing performance, on par with the best and most compute intensive supervised metric.

The most notable limitation is that achieving exponential scaling requires a high quality data pruning metric. Since most metrics developed for smaller datasets scale poorly to ImageNet, our results emphasize the importance of future work in identifying high quality, scalable metrics. Our self-supervised metric provides a strong initial baseline. Moreover, a key advantage of data pruning is reduced computational cost due to training on a smaller dataset for the same number of epochs as the full dataset (see App. C). However, we found that performance often increased when training on the pruned dataset for the same number of iterations as on the full dataset, resulting in the same training time, but additional training epochs. However, this performance gain saturated before training time on the pruned dataset approached that on the whole dataset (App. J) thereby still yielding a computational efficiency gain. Overall this tradeoff between accuracy and training time on pruned data is important to consider in evaluating potential gains due to data pruning. Finally, we found that class-balancing was essential to maintain performance on data subsets (App. H). Future work will be required to identify ways to effectively select the appropriate amount of class-balancing.

A potential negative societal impact could be that data-pruning leads to unfair outcomes for certain groups. We have done a preliminary analysis of how data-pruning affects performance on individual ImageNet classes (App. I), finding no substantial differential effects across classes. However proper fairness tests specific to deployment settings should always be conducted on every model, whether trained on pruned data or not. Additionally, we analyzed the impact of pruning on OOD performance (App. K).

We believe the most promising future direction is the further development of scalable, unsupervised data pruning metrics. Indeed our theory predicts that the application of pruning metrics on larger scale datasets should yield larger gains by allowing more aggressive pruning. This makes data pruning especially exciting for use on the massive unlabeled datasets used to train large foundation models (e.g. 400M image-text pairs for CLIP , 3.5B Instagram images , 650M images for the DALLE-2 encoder , 780B tokens for PALM ). If highly pruned versions of these datasets can be used to train a large number of different models, one can conceive of such carefully chosen data subsets as foundation datasets in which the initial computational cost of data pruning can be amortized across efficiency gains in training many downstream models, just at the initial computational cost of training foundation models is amortized across the efficiency gains of fine-tuning across many downstream tasks. Together, our results demonstrate the promise and potential of data pruning for large-scale training and pretraining.

We thank Priya Goyal, Berfin Simsek, Pascal Vincent valuable discussions, Qing Jin for insights about the optimal pruning distribution, Isaac Seessel for VISSL support as well as Kashyap Chitta and José Álvarez for kindly providing their ensemble active learning score.

References

Appendix

All code required to reproduce the theory figures and numerical simulations throughout this paper can be run in the Colab notebook at https://colab.research.google.com/drive/1in35C6jh7y_ynwuWLBmGOWAgmUgpl8dF?usp=sharing.

In particular, consider pruning the training dataset by keeping only the examples with the smallest margin ∣zμ∣=∣Jprobe⋅xμ∣|z^{\mu}|=|\textbf{J}_{\text{probe}}\cdot\textbf{x}^{\mu}| along a probe student Jprobe\textbf{J}_{\text{probe}}. The pruned dataset will follow some distribution p(z)p(z) along the direction of Jprobe\textbf{J}_{\text{probe}}, and remain isotropic in the nullspace of Jprobe\textbf{J}_{\text{probe}}. In what follows we will derive a general theory for an arbitrary data distribution p(z)p(z), and specialize to the case of small-margin pruning only at the very end (in which case p(z)p(z) will take the form of a truncated Gaussian). We will also make no assumptions on the form of the probe student Jprobe\textbf{J}_{\text{probe}} or the learning rule used to train it; only that Jprobe\textbf{J}_{\text{probe}} has developed some overlap with the teacher, quantified by the angle \theta=\cos^{-1}\big{(}\frac{\textbf{J}_{\text{probe}}\cdot\textbf{T}}{\|\textbf{J}_{\text{probe}}\|_{2}\|\textbf{T}\|_{2}}\big{)} (Fig. 2A).

After the dataset has been pruned, we consider training a new student JJ from scratch on the pruned dataset. A typical training algorithm (used in support vector machines and the solution to which SGD converges on separable data) is to find the solution JJ which classifies the training data with the maximal margin κ=min⁡μJ⋅(yμxμ)\kappa=\min_{\mu}\textbf{J}\cdot(y^{\mu}\textbf{x}^{\mu}). Our goal is to compute the generalization error εg\varepsilon_{g} of this student, which is simply governed by the overlap between the student and the teacher, εg=cos⁡−1(R)/π\varepsilon_{g}=\cos^{-1}(R)/\pi, where R=J⋅T/∥J∥2∥T∥2R=\textbf{J}\cdot\textbf{T}/\|\textbf{J}\|_{2}\|\textbf{T}\|_{2}.

A.2 Main result and overview

Our main result is a set of self-consistent equations which can be solved to obtain the generalization error ε(α,p,θ)\varepsilon(\alpha,p,\theta) for any α\alpha and any data distribution p(z)p(z) along a probe student at any angle θ\theta relative to the teacher. These equations take the form,

R−ρcos⁡θsin⁡2θ\displaystyle\frac{R-\rho\cos\theta}{\sin^{2}\theta} \displaystyle=\frac{\alpha}{\pi\Lambda}\bigg{<}\int_{-\infty}^{\kappa}dt\ \exp\left(-\frac{\Delta(t,z)}{2\Lambda^{2}}\right)(\kappa-t)\bigg{>}_{z} (1) 1−ρ2+R2−2ρRcos⁡θsin⁡2θ\displaystyle 1-\frac{\rho^{2}+R^{2}-2\rho R\cos\theta}{\sin^{2}\theta} \displaystyle=2\alpha\bigg{<}\int_{-\infty}^{\kappa}dt\frac{e^{-\frac{(t-\rho z)^{2}}{2(1-\rho^{2})}}}{\sqrt{2\pi}\sqrt{1-\rho^{2}}}H\bigg{(}\frac{\Gamma(t,z)}{\sqrt{1-\rho^{2}}\Lambda}\bigg{)}(\kappa-t)^{2}\bigg{>}_{z} (2) ρ−Rcos⁡θsin⁡2θ\displaystyle\frac{\rho-R\cos\theta}{\sin^{2}\theta} \displaystyle=2\alpha\bigg{<}\int_{-\infty}^{\kappa}dt\frac{e^{-\frac{(t-\rho z)^{2}}{2(1-\rho^{2})}}}{\sqrt{2\pi}\sqrt{1-\rho^{2}}}H\bigg{(}\frac{\Gamma(t,z)}{\sqrt{1-\rho^{2}}\Lambda}\bigg{)}\bigg{(}\frac{z-\rho t}{1-\rho^{2}}\bigg{)}(\kappa-t) \displaystyle\quad\quad\quad+\frac{1}{2\pi\Lambda}\exp\left(-\frac{\Delta(t,z)}{2\Lambda^{2}}\right)\bigg{(}\frac{\rho R-\cos\theta}{1-\rho^{2}}\bigg{)}(\kappa-t)\bigg{>}_{z} (3) Where, Λ\displaystyle\Lambda =sin⁡2θ−R2−ρ2+2ρRcos⁡θ,\displaystyle=\sqrt{\sin^{2}\theta-R^{2}-\rho^{2}+2\rho R\cos\theta}, (4) Γ(t,z)\displaystyle\Gamma(t,z) =z(ρR−cos⁡θ)−t(R−ρcos⁡θ),\displaystyle=z(\rho R-\cos\theta)-t(R-\rho\cos\theta), (5) Δ(t,z)\displaystyle\Delta(t,z) =z2(ρ2+cos⁡2θ−2ρRcos⁡θ)+2tz(Rcos⁡θ−ρ)+t2sin⁡2θ.\displaystyle=z^{2}\left(\rho^{2}+\cos^{2}\theta-2\rho R\cos\theta\right)+2tz(R\cos\theta-\rho)+t^{2}\sin^{2}\theta. (6)

Where ⟨⋅⟩z\langle\cdot\rangle_{z} represents an average over the pruned data distribution p(z)p(z) along the probe student. For any α,p(z),θ\alpha,p(z),\theta, these equations can be solved for the order parameters R,ρ,κR,\rho,\kappa, from which the generalization error can be easily read off as εg=cos⁡−1(R)/π\varepsilon_{g}=\cos^{-1}(R)/\pi. This calculation results in the solid theory curves in Figs 1,2,3, which show an excellent match to numerical simulations. In the following section we will walk through the derivation of these equations using replica theory. In Section A.6 we will derive an expression for the information gained per training example, and show that with Pareto optimal data pruning this information gain can be made to converge to a finite rate, resulting in at least exponential decay in test error. In Section A.7, we will show that super-exponential scaling eventually breaks down when the probe student does not match the teacher perfectly, resulting in power law scaling at at a minimum pruning fraction fmin(θ).f_{\text{min}}(\theta).

A.3 Replica calculation of the generalization error

To obtain Eqs. 1,2,3, we follow the approach of Elizabeth Gardner and compute the volume Ω(xμ,T,κ)\Omega(\textbf{x}^{\mu},\textbf{T},\kappa) of solutions JJ which perfectly classify the training data up to a margin κ\kappa (known as the Gardner volume) . As κ\kappa grows, the volume of solutions shrinks until it reaches a unique solution at a critical κ\kappa, the max-margin solution. The Gardner volume Ω\Omega takes the form,

Because the student’s decision boundary is invariant to an overall scaling of J, we enforce normalization of J via the measure dμ(J)d\mu(\textbf{J}),

In the thermodynamic limit N,P→∞N,P\to\infty the typical value of the entropy S(κ)=⟨⟨log⁡Ω(xμ,T,κ)⟩⟩S(\kappa)=\langle\langle\log\Omega(\textbf{x}^{\mu},\textbf{T},\kappa)\rangle\rangle is dominated by particular values of R,κR,\kappa, where the double angle brackets ⟨⟨⋅⟩⟩\langle\langle\cdot\rangle\rangle denote a quenched average over disorder introduced by random realizations of the training examples xμ\textbf{x}^{\mu} and the teacher T. However, computing this quenched average is intractable since the integral over J cannot be performed analytically for every individual realization of the examples. We rely on the replica trick from statistical physics,

Which allows us to evaluate S(κ)S(\kappa) in terms of easier-to-compute powers of Ω\Omega,

This reduces our problem to computing powers of Ω\Omega, which for integer nn can be written in terms of α=1,…,n\alpha=1,\ldots,n replicated copies of the original system,

We begin by introducing auxiliary variables,

by δ−\delta-functions, to pull the dependence on J and T outside of the Heaviside function,

Using the integral representation of the δ\delta-functions,

The data obeys some distribution p(z)p(z) along the direction of Jprobe\textbf{J}_{\text{probe}} and is isotropic in the nullspace of Jprobe\textbf{J}_{\text{probe}}. Hence we can decompose a training example xμ\textbf{x}^{\mu} as follows, xμ=Jprobezμ+(I−JprobeJprobeT)sμ\textbf{x}^{\mu}=\textbf{J}_{\text{probe}}z^{\mu}+(I-\textbf{J}_{\text{probe}}\textbf{J}_{\text{probe}}^{T})\textbf{s}^{\mu}, where zμ∼p(z)z^{\mu}\sim p(z) and sμ∼N(0,IN)\textbf{s}^{\mu}\sim\mathcal{N}(0,I_{N}),

Where J⊥=(1−JprobeJprobeT)J\textbf{J}_{\perp}=(1-\textbf{J}_{\text{probe}}\textbf{J}_{\text{probe}}^{T})\textbf{J} and T⊥=(1−JprobeJprobeT)T\textbf{T}_{\perp}=(1-\textbf{J}_{\text{probe}}\textbf{J}_{\text{probe}}^{T})\textbf{T}. Now we can average over the patterns sμ∼N(0,IN)\textbf{s}^{\mu}\sim\mathcal{N}(0,I_{N}),

Inserting this back into our expression for the Gardner volume,

As is typical in replica calculations of this type, we now introduce order parameters,

which will allow us to decouple the J- from the λ\lambda-μ\mu-zz- integrals. qαβq^{\alpha\beta} represents the overlaps between replicated students, and RαR^{\alpha} the overlap between each replicated student and the teacher. However, because our problem involves the additional role of the probe student, we must introduce an additional order parameter,

which represents the overlap between each replicated student and the probe student. Notice that,

With this new set of order parameters in hand, we can decouple the J from the λ−u−z−\lambda-u-z-integrals.

We can now perform the gaussian integral over u^μ\hat{u}_{\mu},

Now we introduce integral representations for the remaining delta functions, including the measure dμ(Jα)d\mu(\textbf{J}^{\alpha}), for which we introduce the parameter k^α\hat{k}^{\alpha},

Notice that the uμ−λμα−λ^μα−zμu_{\mu}-\lambda_{\mu}^{\alpha}-\hat{\lambda}_{\mu}^{\alpha}-z_{\mu}-integrals factorize in μ\mu, and can be written as a single integral to the power of P=αNP=\alpha N.

Where we have written the Gardner volume in terms of an entropic part GSG_{S}, which measures how many spherical couplings satisfy the constraints,

We first evaluate the entropic part, GSG_{S}, by introducing the n×nn\times n matrices A,BA,B,

Inserting this our expression for GSG_{S} becomes

Now we can include the remaining terms in the expression for Ω(n)\Omega^{(n)} outside of GEG_{E} and GSG_{S} by noting that

Additionally, we can use log⁡det⁡A=tr(log⁡A)\log\det A=tr(\log A). Thus all terms in the exponent except GEG_{E} can be written as

Now we extremize wrt R^α\hat{R}^{\alpha} and the elements of AA by setting the derivatives wrt R^γ\hat{R}^{\gamma}, ρ^γ\hat{\rho}^{\gamma} and AγδA^{\gamma\delta} equal to zero:

In order to extremize wrt qαβ,Rα,ραq^{\alpha\beta},R^{\alpha},\rho^{\alpha}, we take the replica symmetry ansatz ,

A matrix with EE on the diagonal and FF elsewhere, Cαβ=Eδαβ+F(1−δαβ)C_{\alpha\beta}=E\delta_{\alpha\beta}+F(1-\delta_{\alpha\beta}), has n−1n-1 degenerate eigenvalues E−FE-F and one eigenvalue E+(n−1)FE+(n-1)F. Hence in our case CC has n−1n-1 degenerate eigenvalues

We next evaluate the energetic part, GEG_{E},

To simplify the last term we apply the Hubbard-Stratonovich transformation, eb2/2=∫Dtebte^{b^{2}/2}=\int Dte^{bt}, introducing auxiliary field tt,

Using the Θ\Theta-function to restrict the bounds of integration,

Now we can perform the gaussian integrals over λ^α\hat{\lambda}^{\alpha},

Shifting the integration variable t→(ραz+(u−zcos⁡θ)(Rα−ραcos⁡θ)sin⁡2θ+q−ρ2−(R−ρcos⁡θ)2sin⁡2θt)/qt\to(\rho^{\alpha}z+\frac{(u-z\cos\theta)(R^{\alpha}-\rho^{\alpha}\cos\theta)}{\sin^{2}\theta}+\sqrt{q-\rho^{2}-\frac{(R-\rho\cos\theta)^{2}}{\sin^{2}\theta}}t)/\sqrt{q}, we can finally perform the gaussian integral over uu,

We can simplify this further by taking t→(qt−zρ)/q−ρ2t\to(\sqrt{q}t-z\rho)/\sqrt{q-\rho^{2}},

A.4 Quenched entropy

Putting everything together, and using the replica identity, \langle\log\text{\Omega\rangle==\lim_{n\to 0}$(⟨Ωn⟩−1)(\langle\Omega^{n}\rangle-1)}/n,$ we obtain an expression for the quenched entropy of the teacher-student perceptron under data pruning:

\frac{1}{N}\langle\log\Omega\rangle=\text{extr}_{q,R,\rho}\bigg{[}\frac{1}{2}\log\bigg{(}1-q\bigg{)}+\frac{1}{2}\bigg{(}\frac{q-(R^{2}-2R\rho\cos\theta+\rho^{2})/\sin^{2}\theta}{1-q}\bigg{)}\\ +2\alpha\bigg{<}\int Dt\log H\bigg{(}\frac{\kappa-(z\rho+\sqrt{q-\rho^{2}})t}{\sqrt{1-q}}\bigg{)}\\ \times H\bigg{(}\frac{t\left(\sqrt{q-\rho^{2}}+\rho z\right)(R-\rho\cos\theta)+z(q\cos\theta-\rho R)}{\sqrt{\left(q-\rho^{2}\right)\left(R^{2}+\rho^{2}-q\sin^{2}\theta-2\rho R\cos\theta\right)}}\bigg{)}\bigg{>}_{z}\bigg{]} (63)

We will now unpack this equation and use it to make predictions in several specific settings.

A.5 Perfect teacher-probe overlap

We will begin by considering the case where the probe student has learned to perfectly match the teacher, Jprobe=TJ_{\text{probe}}=T, which we can obtain by the limit θ→0\theta\to 0, ρ→R\rho\to R. In this limit the second HH-function in Eq. 63 becomes increasingly sharp, approaching a step function:

\frac{1}{N}\langle\langle\ln\Omega(\textbf{x}^{\mu},T,\kappa)\rangle\rangle=\text{extr}_{q,R}\bigg{[}\frac{1}{2}\log\bigg{(}1-q\bigg{)}+\frac{1}{2}\bigg{(}\frac{q-R^{2}}{1-q}\bigg{)}\\ +2\alpha\int Dt\int dzp(z)\Theta(z)\log H\bigg{(}-\frac{\sqrt{q-R^{2}}t+Rz-\kappa}{\sqrt{1-q}}\bigg{)}\bigg{]} (65)

We can now obtain a set of self-consistent saddle point equations by setting set to zero the derivatives with respect to RR and qq of the right side of Eq. 65. As κ\kappa approaches its critical value, the space of solutions shrinks to a unique solution, and hence the overlap between students qq approaches one. In the limit q→1q\to 1, after some partial integration, we find,

R=2\alpha\int_{-\infty}^{\kappa}\frac{dt}{\sqrt{2\pi}\sqrt{1-R^{2}}}\int_{0}^{\infty}dzp(z)\exp\bigg{(}-\frac{(t-Rz)^{2}}{2(1-R^{2})}\bigg{)}\bigg{(}\frac{z-Rt}{1-R^{2}}\bigg{)}(\kappa-t) (66) 1-R^{2}=2\alpha\int_{-\infty}^{\kappa}\frac{dt}{\sqrt{2\pi}\sqrt{1-R^{2}}}\int_{0}^{\infty}dzp(z)\exp\bigg{(}-\frac{(t-Rz)^{2}}{2(1-R^{2})}\bigg{)}\big{(}\kappa-t\big{)}^{2} (67)

These saddle point equations can be solved numerically to find RR and κ\kappa as a function of α\alpha for a student perceptron trained on a dataset with an arbitrary distribution along the teacher direction p(z)p(z). We can specialize to the case of data pruning by setting p(z)p(z) to the distribution found after pruning an initially Gaussian-distributed dataset so that only a fraction ff of those examples with the smallest margin along the teacher are kept, p(z)=e−z2/22πfΘ(γ−∣z∣)p(z)=\frac{e^{-z^{2}/2}}{\sqrt{2\pi}f}\Theta(\gamma-|z|), where the threshold \gamma=H^{-1}\big{(}\frac{1-f}{2}\big{)}.

R=\frac{2\alpha}{f\sqrt{2\pi}\sqrt{1-R^{2}}}\int_{-\infty}^{\kappa}Dt\ \exp\bigg{(}-\frac{R^{2}t^{2}}{2(1-R^{2})}\bigg{)}\bigg{[}1-\exp\bigg{(}-\frac{\gamma(\gamma-2Rt)}{2(1-R^{2})}\bigg{)}\bigg{]}(\kappa-t) (68) 1-R^{2}=\frac{2\alpha}{f}\int_{-\infty}^{\kappa}Dt\ \bigg{[}H\bigg{(}-\frac{Rt}{\sqrt{1-R^{2}}}\bigg{)}-H\bigg{(}-\frac{Rt-\gamma}{\sqrt{1-R^{2}}}\bigg{)}\bigg{]}(\kappa-t)^{2} (69)

Solving these saddle point equations numerically for RR and κ\kappa yields an excellent fit to numerical simulations, as can be seen in Fig. 1A. It is also easy to verify that in the limit of no data pruning (f→1,γ→∞f\to 1,\gamma\to\infty) we recover the saddle point equations for the classical teacher-student perceptron (Eqs. 4.4 and 4.5 in ),

A.6 Information gain per example

Why does data pruning allow for super-exponential performance with dataset size α\alpha? We can define the amount of information gained from each new example, I(α),I(\alpha),as the fraction by which the space of solutions which perfectly classify the data is reduced when a new training example is added, I(α)=Ω(P+1N)/Ω(PN)I(\alpha)=\Omega(\frac{P+1}{N})/\Omega(\frac{P}{N}). Or, equivalently, the rate at which the entropy is reduced, I(α)=−ddαS(α)I(\alpha)=-\frac{d}{d\alpha}S(\alpha). Of coure, the volume of solutions shrinks to zero at the max-margin solution; so to study the volume of solutions which perfectly classify the data we simply set the margin to zero κ=0\kappa=0. In the information gain for a perceptron in the classical teacher-student setting is shown to take the form,

Which goes to zero in the limit of large α\alpha as I(α)∼1/αI(\alpha)\sim 1/\alpha. Data pruning can increase the information gained per example by pruning away the uninformative examples. To show this, we generalize the calculation of the information gain to pruned datasets, using the expression for the entropy we obtained in the previous section (Eq. 65).

Hence the information gain I(α)=−ddαS(α)I(\alpha)=-\frac{d}{d\alpha}S(\alpha) is given by

Changing variables to t→−(Rt+R1−Rz)/qt\to-(\sqrt{R}t+\frac{R}{\sqrt{1-R}}z)/\sqrt{q},

Now, assuming that we prune to a fraction ff, so that p(z)=\Theta\big{(}|z|-\gamma\big{)}\frac{\exp(-z^{2}/2)}{\sqrt{2\pi}f}, where \gamma=H^{-1}\bigg{(}\frac{1-f}{2}\bigg{)}

I(α)I(\alpha) is plotted for varying values of ff in Fig. 1F. Notice that for f→1f\to 1, γ→∞\gamma\to\infty and we recover Eq. 72. To obtain the optimal pruning fraction foptf_{\text{opt}} for any α\alpha, we first need an equation for RR, which can be obtained by taking the saddle point of Eq. 73. Next we optimize I(α)I(\alpha) by setting the derivative of Eq. 76 with respect to ff equal to zero. This gives us a pair of equations which can be solved numerically to obtain foptf_{\text{opt}} for any α\alpha.

Finally, Eq. 76 reveals that as we prune more aggressively the information gain per example approaches a finite rate. As f→0f\to 0, γ→0\gamma\to 0, and we obtain,

Which allows us to produce to trace the Pareto frontier in Fig. 1F. For R→1R\to 1, Eq. 77 gives the asymptotic information gain I(∞)=1I(\infty)=1 nat/example.

A.7 Imperfect teacher-probe overlap

In realistic settings we expect the probe student to have only partial information about the target function. What happens if the probe student does not perfectly match the teacher? To understand this carefully, we need to compute the full set of saddle point equations over R,q,R,q, and ρ\rho, which we will do in the following section. But to first get an idea for what goes wrong, we include in this section a simple sketch which reveals the limiting behavior.

Consider the case where the angle between the probe student and teacher is θ\theta. Rotate coordinates so that the first canonical basis vector aligns with the student J=(1,0,…,0)J=(1,0,\ldots,0), and the teacher lies in the span of the first two canonical basis vectors, T=(cos⁡θ,sin⁡θ,0,…,0)T=(\cos\theta,\sin\theta,0,\ldots,0). Consider the margin along the teacher of a new training example xx drawn from the pruned distribution.

Hence the data ultimately stops concentrating around the teacher’s decision boundary, and the information gained from each new example goes to zero. Therefore we expect the generalization error to converge to a power law, where the constant prefactor is roughly that of pruning with a prune fraction fminf_{\text{min}} which yields an average margin of 1−R21-R^{2}. This “minimum” pruning fraction lower bounds the generalization error envelope (see Fig. 2), and satisfies the following equation,

where \gamma_{\text{min}}=H^{-1}\bigg{(}\frac{1-f_{\text{min}}}{2}\bigg{)}. Eq. 80 can be solved numerically, and we use it to produce the lower-bounding power laws shown in red in Fig. 2C,D. The minimum achievable pruning fraction fmin(θ)f_{\text{min}}(\theta) approaches zero as the angle between the probe student and the teacher shrinks, and we can obtain its scaling by taking R→1R\to 1, in which case we find,

A.8 Optimal pruning policy

The saddle point equations Eq. 66,67 reveal that the optimal pruning policy varies as a function of αprune\alpha_{\text{prune}}. For αprune\alpha_{\text{prune}} large the best policy is to retain only the “hardest" (smallest-margin) examples. But when αprune\alpha_{\text{prune}} is small, keeping the “hardest" examples performs worse than chance, suggesting that the best policy in the αprune\alpha_{\text{prune}} small regime is to keep the easiest examples. Indeed by switching between the “keep easy” and “keep hard” strategies as αprune\alpha_{\text{prune}} grows, one can achieve a lower Pareto frontier than the one shown in Fig. 1A in the small αprune\alpha_{\text{prune}} regime (Fig. 7C).

These observations beg the question: what is the best policy in the intermediate αprune\alpha_{\text{prune}} regime? Is there a globally optimal pruning policy which interpolates between the “keep easy” and “keep hard” strategies and achieves the lowest possible Pareto frontier (blue curve in Fig. 7A)?

In this section we investigate this question. Using the calculus of variations, we first derive the optimal data distribution p(z∣αprune,f)p(z|\alpha_{\text{prune}},f) along the teacher for a given αprune\alpha_{\text{prune}}, ff. We begin by framing the problem using the method of Lagrange multipliers. Seeking to optimize RR under the constraints imposed by the saddle point equations Eqs. 66,67, we define the Lagrangian,

Taking a variational derivative δLδp\frac{\delta\mathcal{L}}{\delta p} with respect to the data distribution pp, we obtain an equation for zz, indicating that the optimal distribution is a delta function at z=z∗z=z^{*}. To find the optimal location of the delta function z∗z^{*}, we take derivatives with respect to the remaining variables R,k,μ,λR,k,\mu,\lambda and solve the resulting set of equations numerically. The qualitative behavior is shown in Fig. 7A. As αprune\alpha_{\text{prune}} grows, the location of the delta function shifts from infinity to zero, confirming that the optimal strategy for small αprune\alpha_{\text{prune}} is to keep the "easy" (large-margin) examples, and for large αprune\alpha_{\text{prune}} to keep the "hard" (small-margin) examples.

Interestingly, this calculation also reveals that if the location of the delta function is chosen optimally, the student can perfectly recover the teacher (R=1R=1, zero generalization error) for any αprune\alpha_{\text{prune}}. This observation, while interesting, is of no practical consequence because it relies on an infinitely large training set from which examples can be precisely selected to perfectly recover the teacher. Therefore, to derive the optimal pruning policy for a more realistic scenario, we assume a gaussian distribution of data along the teacher direction and model pruning as keeping only those examples which fall inside a window a<z<ba<z<b. The saddle point equations, Eqs. 66,67, then take the form,

Where aa must satisfy a=H−1(f/2+H(b))a=H^{-1}(f/2+H(b)). For each f,αf,\alpha, we find the optimal location of this window using the method of Lagrange multipliers. Defining the Lagrangian as before,

To find the optimal location of the pruning window, we take derivatives with respect to the remaining variables b,R,k,μ,λb,R,k,\mu,\lambda and solve the resulting set of equations numerically. Consistent with the results for the optimal distribution, the location of the optimal window shifts from around infinity to around zero as αprune\alpha_{\text{prune}} grows (Fig. 7C).

A.9 Exact saddle point equations

To obtain exact expressions for the generalization error for all θ\theta, we can extremize Eq. 63 wrt R,q,ρR,q,\rho.

Integrating the right-hand side by parts,

Changing variables to t→tq−ρ2+ρzt\to t\sqrt{q-\rho^{2}}+\rho z and taking the limit q→1q\to 1,

Where Γ(t,z)=(R−ρcos⁡θ)(tq−ρ2+ρz)+qzcos⁡θ−ρRz\Gamma(t,z)=(R-\rho\cos\theta)\left(t\sqrt{q-\rho^{2}}+\rho z\right)+qz\cos\theta-\rho Rz. After integating by parts,

Changing variables to t→tq−ρ2+ρzt\to t\sqrt{q-\rho^{2}}+\rho z and taking the limit q→1q\to 1,

Where now Γ(t,z)=z(ρR−cos⁡θ)−t(R−ρcos⁡θ).\Gamma(t,z)=z(\rho R-\cos\theta)-t(R-\rho\cos\theta).

Changing variables to t→tq−ρ2+ρzt\to t\sqrt{q-\rho^{2}}+\rho z and taking the limit q→1q\to 1,

So together we have three saddle point equations:

R−ρcos⁡θsin⁡2θ\displaystyle\frac{R-\rho\cos\theta}{\sin^{2}\theta} \displaystyle=\frac{\alpha}{\pi\Lambda}\bigg{<}\int_{-\infty}^{\kappa}dt\ \exp\left(-\frac{\Delta(t,z)}{2\Lambda^{2}}\right)(\kappa-t)\bigg{>}_{z} (115) 1−ρ2+R2−2ρRcos⁡θsin⁡2θ\displaystyle 1-\frac{\rho^{2}+R^{2}-2\rho R\cos\theta}{\sin^{2}\theta} \displaystyle=2\alpha\bigg{<}\int_{-\infty}^{\kappa}dt\frac{e^{-\frac{(t-\rho z)^{2}}{2(1-\rho^{2})}}}{\sqrt{2\pi}\sqrt{1-\rho^{2}}}H\bigg{(}\frac{\Gamma(t,z)}{\sqrt{1-\rho^{2}}\Lambda}\bigg{)}(\kappa-t)^{2}\bigg{>}_{z} (116) ρ−Rcos⁡θsin⁡2θ\displaystyle\frac{\rho-R\cos\theta}{\sin^{2}\theta} \displaystyle=2\alpha\bigg{<}\int_{-\infty}^{\kappa}dt\frac{e^{-\frac{(t-\rho z)^{2}}{2(1-\rho^{2})}}}{\sqrt{2\pi}\sqrt{1-\rho^{2}}}H\bigg{(}\frac{\Gamma(t,z)}{\sqrt{1-\rho^{2}}\Lambda}\bigg{)}\bigg{(}\frac{z-\rho t}{1-\rho^{2}}\bigg{)}(\kappa-t) (117) \displaystyle\quad\quad\quad+\frac{1}{2\pi\Lambda}\exp\left(-\frac{\Delta(t,z)}{2\Lambda^{2}}\right)\bigg{(}\frac{\rho R-\cos\theta}{1-\rho^{2}}\bigg{)}(\kappa-t)\bigg{>}_{z} (118) Where Λ\displaystyle\Lambda =sin⁡2θ−R2−ρ2+2ρRcos⁡θ,\displaystyle=\sqrt{\sin^{2}\theta-R^{2}-\rho^{2}+2\rho R\cos\theta}, (119) Γ(t,z)\displaystyle\Gamma(t,z) =z(ρR−cos⁡θ)−t(R−ρcos⁡θ),\displaystyle=z(\rho R-\cos\theta)-t(R-\rho\cos\theta), (120) Δ(t,z)\displaystyle\Delta(t,z) =z2(ρ2+cos⁡2θ−2ρRcos⁡θ)+2tz(Rcos⁡θ−ρ)+t2sin⁡2θ.\displaystyle=z^{2}\left(\rho^{2}+\cos^{2}\theta-2\rho R\cos\theta\right)+2tz(R\cos\theta-\rho)+t^{2}\sin^{2}\theta. (121)

Solving these equations numerically yields an excellent fit to numerical simulations on structured data (Fig. 2BCD).

Appendix B Model training method details & dataset information

ImageNet model training was performed using a standard ResNet-50 through the VISSL library (stable version v0.1.6), which provides default configuration files for supervised ResNet-50 training (accessible here; released under the MIT license). Each model was trained on a single node of 8 NVIDIA V100 32GB graphics cards with BATCHSIZE_PER_REPLICA = 256, using the Stochastic Gradient Descent (SGD) optimizer with a base learning rate = 0.1, nesterov momentum = 0.9, and weight decay = 0.001. For our scaling experiments (Fig. 3 and Fig. 9), we trained one model per fraction of data kept (0.1-1.0) for each dataset size. In total, these plot required training 97 models on (potentially a subset of) ImageNet. All the models were trained with matched number of iterations, corresponding to 105 epochs on the full ImageNet dataset. The learning rate was decayed by a factor of 10 after the number of iterations corresponding to 30, 60, 90, and 100 epochs on the full ImageNet dataset.

For our main ImageNet experiments (Fig. 5) we trained one model per fraction of data kept (1.0, 0.9, 0.8, 0.7, 0.6) ×\times metric (11 metrics in total). In the plot itself, since any variation in the “fraction of data kept = 1.0” setting is due to random variation across runs not due to potential metric differences, we averaged model performances to obtain a single datapoint here (while also keeping track of the variation across models, which is plotted as ±2\pm 2 standard deviations). In total, this plot required training 55 models on (potentially a subset of) ImageNet. For Fig. 5C, in order to reduce noise from random variation, we additionally trained five models per datapoint and metric, and plot the averaged performance in addition to error bars showing one standard deviation of the mean. Numerical results from Figure 5BC are available from Table 1. ImageNet is released under the ImageNet terms of access. It is important to note that ImageNet images are often biased . The SWaV model used to compute our prototypicality metrics was obtained via torch.hub.load(‘facebookresearch/swav:main’, ‘resnet50’), which is the original model provided by ; we then used the avgpool layer’s activations.

CIFAR-10 and SVHN model training was performed using a standard ResNet-18 through the PyTorch library. Each model was trained on a single NVIDIA TITAN Xp 12GB graphics card with batch size = 128, using the Stochastic Gradient Descent (SGD) optimizer with learning rate = 0.1, nesterov momentum = 0.9, and weight decay = 0.0005. Probe models were trained for 20 epochs each for CIFAR-10 and 40 epochs each for SVHN. Pruning scores were then computed using the EL2Ns metric , averaged across 10 independent initializations of the probe models. To evaluate data pruning performance, fresh models were trained from scratch on each pruned dataset for 200 epochs, with the learning rate decayed by a factor of 5 after 60, 120 and 160 epochs.

To assess the effect of pruning downstream finetuning data on transfer learning performance, vision transformers (ViTs) pre-trained on ImageNet21k were fine-tuned on different pruned subsets of CIFAR-10. Pre-trained models were obtained from the timm model library . Each model was trained on a single NVIDIA TITAN Xp 12GB graphics card with batch size = 128, using the Adam optimizer with learning rate = 1e-5 and no weight decay. Probe models were trained for 2 epochs each. Pruning scores were then computed using the EL2Ns metric , averaged across 10 independent random seeds. To evaluate data pruning performance, pre-trained models were fine-tuned on each pruned dataset for 10 epochs.

To assess the effect of pruning upstream pretraining data on transfer learning performance, each of the ResNet-50s pre-trained on pruned subsets of ImageNet1k in Fig. 3D was fine-tuned on all of CIFAR-10. Each model was trained on a single NVIDIA TITAN Xp 12GB graphics card with batch size = 128, using the RMSProp optimizer with learning rate = 1e-4 and no weight decay. Probe models were trained for 2 epochs each. Pruning scores were then computed using the EL2Ns metric , averaged across 10 independent random seeds. To evaluate data pruning performance, pre-trained models were fine-tuned on each pruned dataset for 10 epochs.

Appendix C Breaking compute scaling laws via data pruning

Do the savings in training dataset size we have identified translate to savings in compute, and can data pruning be used to beat widely observed compute scaling laws ? Here we show for the perceptron that data pruning can afford exponential savings in compute, and we provide preliminary evidence that the same is true for ResNets trained on CIFAR-10 and ImageNet. We repeat the perceptron learning experiments in Fig. 1A, keeping track of the computational complexity of each experiment, measured by the time to convergence of the quadratic programming algorithm used to find a max-margin solution (see B for details). Across all experiments, the convergence time TT was linearly proportional to αprune\alpha_{\text{prune}} with T=0.96αprune+0.80T=0.96\alpha_{\text{prune}}+0.80, allowing us to replace the x-axis of 1A with compute to produce Fig. 8A, which reveals that data pruning can be used to break compute scaling laws for the perceptron.

Motivated by this, we next investigate whether the convergence time of neural networks trained on pruned datasets depends largely on the number of examples and not their difficulty, potentially allowing for exponential compute savings. We investigate the learning curves of a ResNet18 trained on CIFAR-10 and a ResNet50 on ImageNet for several different pruning fractions (Fig. 8B). While previous works have fixed the number of iterations , here we fix the number of epochs, so that the model trained on 60% of the full dataset is trained for only 60% the iterations of the model trained on the full dataset, using only 60% the compute. Nevertheless, we find that the learning curves are strikingly similar across pruning fractions, and appear to converge equally quickly. These results suggest that data pruning could lead to large compute savings in practical settings, and in ongoing experiments we are working to make the analogs of Fig. 8A for ResNets on CIFAR-10 and ImageNet to quantify this benefit.

Appendix D Additional scaling experiments

In Fig. 9 we perform additional scaling experiments using the EL2Ns and self-supervised prototypes metrics. In Fig. 10 we give a practical example of a cross over from exponential to power-law scaling when the probe student has limited information about the teacher (here a model trained for only a small number of epochs on SVHN) .

Appendix E Extremal images according to different metrics

In Fig. 6, we showed extremal images for two metrics (self-supervised prototypes, memorization) and a single class. In order to gain a better understanding of how extremal images (i.e. images that are easiest or hardest to learn according to different metrics) look like for all metrics and more classes, we here provide additional figures. In order to avoid cherry-picking classes while at the same time making sure that we are visualizing images for very different classes, we here show extreme images for classes 100, …, 500 while leaving out classes 0 and 400 (which would have been part of the visualization) since those classes almost exclusively consist of images containing people (0: tench, 400: academic gown). The extremal images are shown in Figures 12,13,14,15,16,17,18,19.

Appendix F Impact of number of clusters k𝑘k on self-supervised prototypes

Our self-supervised prototype metric is based on kk-means clustering, which has a single hyperparameter kk. By default and throughout the main paper, we set k=1000k=1000, corresponding to the number of classes in ImageNet. Here, we investigate other settings of kk to understand how this hyperparameter impacts performance. As can be seen in Table 2, kk does indeed have an impact on performance, and very small values for kk (e.g. k<10k<10) as well as very large values for kk (e.g. k=50,000k=50,000) both lead to performance impairments. At the same time, performance is relatively high across very different in-between settings for kk. In order to assess these results, it may be important to keep in mind that ±0.54%\pm 0.54\% corresponds to plus/minus 2 standard deviations of performance when simply training the same model multiple times (with different random initialization). Overall, these results suggest that if the number of clusters kk deviates at most by one order of magnitude from the number of classes in the dataset (for ImageNet-1K), the exact choice of kk does not matter much.

Appendix G Impact of ensemble prototypes

The self-supervised prototypes metric is based on kk-means clustering in the embedding space of a self-supervised (=SSL) model. Since even otherwise identical models trained with different random seeds can end up with somewhat different embedding spaces, we here investigated how the performance of our self-supervised prototypes metric would change when averaging the scores derived from five models, instead of just using a single model’s score. The results, shown in Table 3, indicate that ensembling the self-supervised prototype scores neither improves nor hurts performance. This is both good and bad news: Bad news since naturally any improvement in metric development leads to better data efficiency; on the other hand this is also good news since ensembles increase the computational cost of deriving the metric—and this suggests that ensembling is not necessary to achieve the performance we achieved (unlike in other methods such as ensemble active learning).

Appendix H Relationship between pruning and class (im-)balance

It is well-known that strong class imbalance in a dataset is a challenge that needs to be addressed. In order to understand the effect of pruning on class (im-)balance, we quantified this relationship. For context, if pruning according to a certain metric preferentially leads to discarding most (or even all) images from certain classes, it is likely that the performance on those classes will drop as a result if this imbalance is not addressed.

Since a standard measure of class (im-)balance—dividing the number of images for the majority class by the number of images for the minority class—is highly sensitive to outliers and discards information about the 998 non-extreme ImageNet classes, we instead calculated a class balance score b∈[0%,100%]b\in[0\%,100\%] as the average class imbalance across any two pairs of classes by computing the expectation over taking two random classes, and then computing how many images the minority class has in proportion to the majority class. For instance, a class balance score of 90% means that on average, when selecting two random classes from the dataset, the smaller of those two classes contains 90% of the number of images of the larger class (higher=better; 100% would be perfectly balanced).

In Fig. 20, we observe that dataset pruning strongly increases class imbalance. This is the case both when pruning away easy images and when pruning away hard images, and the effect occurs for all pruning metrics except, of course, for random pruning. Class imbalance is well-known to be a challenge for deep learning models when not addressed properly . The cause for the amplified class imbalance is revealed when looking at class-conditional differences of metric scores (Figs. 22 and 23): The histograms of the class-conditional score distributions show that for many classes, most (if not all) images have very low scores, while for others most (if not all) images have very high scores. This means that as soon as the lowest / highest scoring images are pruned, certain classes are pruned preferentially and thus class imbalance worsens.

We thus use 50% class balancing for our ImageNet experiments. This ensures that every class has at least 50% of the images that it would have when pruning all classes equally (essentially providing a fixed floor for the minimum number of images per class). This simple fix is an important step to address and counteract class imbalance, although other means could be used as well; and ultimately one would likely want to use a self-supervised version of class (or cluster) balancing when it comes to pruning large-scale unlabeled datasets. For comparison purposes, the results for ImageNet pruning without class balancing are shown in supplementary Fig. 21.

Appendix I Effect of pruning on class-conditional accuracy and fairness

In order to study the effect of dataset pruning on model fairness, at least with respect to specific ImageNet classes, we compared the class-conditional accuracy of a ResNet-50 model trained on the full ImageNet dataset, versus that of the same model trained on an 80% subset obtained after pruning. We used two supervised pruning metrics (EL2N, memorization) and one self-supervised pruning metric (self-supervised prototypes) for obtaining the pruned dataset. In all three cases, and across all 10001000 classes, we found that the class-conditional accuracy of the model trained on a pruned subset of the dataset remains quite similar to that of the model trained on the full dataset (Fig. 24). However, we did notice a very small reduction in class-conditioned accuracy for ImageNet classes that were least accurately predicted by models trained on the entire dataset (blue lines slightly above red unity lines when class-conditioned accuracy is low). This suggests that pruning yields a slight systematic reduction in the accuracy of harder classes, relative to easier classes, though the effect is small.

While we have focused on fairness with respect to individual ImageNet classes, any ultimate test of model fairness should be conducted in scenarios that are specific to the use case of the deployed model. Our examination of the fairness of pruning with respect to individual ImageNet classes constitutes only an initial foray into an exploration of fairness, given the absence of any specific deployment scenario for our current models other than testing them on ImageNet. We leave a full exploration of fairness in other deployment settings for future work.

Appendix J Interaction between data pruning and training duration

Throughout the main paper, our ImageNet experiments are based on a setting where the number of training epochs is kept constant (i.e. we train the same number of epochs on the smaller pruned dataset as on the larger original dataset). This means that data pruning directly reduces the number of iterations required to train the model specifically by reducing the size of the dataset. However, this simultaneously places two constraints on model performance: not only training on a smaller data set, but also training for fewer iterations.

We therefore investigate how model performance changes if we train longer, as quantified by a matched iterations factor. A matched iterations factor of corresponds to the default setting used in the paper of training for the same number of epochs (so that a smaller dataset means proportionally fewer training iterations). In contrast a matched iterations factor of 11 corresponds to training on the smaller pruned dataset for a number of iterations equal to that when training on the larger initial dataset (e.g. when pruning away 50% of the dataset one would train twice as long to match the number of iterations of a model trained on 100% of the dataset). Otherwise the matched iterations factor reflects a linear interpolation in the number of training iterations as the factor varies between and 11.

The results are shown in Table 4 and indicate that training longer does indeed improve performance slightly; however, a matched iterations factor of around 0.4–0.6 may already be sufficient to reap the full benefit. Any matched iterations factor strictly smaller than 1.0 comes with reduced training time compared to training a model on the full dataset.

Appendix K Out-of-distribution (OOD) analysis of dataset pruning

Pruning changes the data regime that a model is exposed to. Therefore, a natural question is how this might affect desirable properties beyond IID performance like fairness (see Appendix I) and out-of-distribution, or OOD, performance which we investigate here. To this end, we use the model-vs-human toolbox based on data and analyses from . This toolbox is comprised of 17 different OOD datasets, including many image distortions and style changes.

In Figure 25(a), OOD accuracies averaged across those 17 datasets are shown for a total of 12 models. These models all have a ResNet-50 architecture . Two baseline models are trained on the full ImageNet training dataset, one using torchvision (purple) and the other using VISSL (blue). Human classification data is shown as an additional reference in red. The remaining 10 models are VISSL-trained on pruned versions of ImageNet using pruning fractions in {0.1, 0.2, 0.3, 0.4, 0.5} and our self-supervised prototype metric. A pruning fraction of 0.3 would correspond to “fraction of data kept = 0.7”, i.e. to training on 70% of ImageNet while discarding the other 30%. We investigated two different settings: discarding easy examples (the default used throughout the paper), which is denoted as “Best Case” (or BC) in the plots; and the reverse setting, i.e. discarding hard examples denoted as Worst Case, or WC. (These terms should be taken with a grain of salt; examples are only insofar best- or worst case examples as predicted by the metric, which itself is in all likelihood far less than perfect.)

The results are as follows: In terms of OOD accuracy (Figure 25(a)), best-case pruning in green achieves very similar accuracies to the most relevant baseline, the blue ResNet-50 model trained via VISSL. This is interesting since oftentimes, OOD accuracies closely follow IID accuracies except for a constant offset , and we know from Figure 5 that the self-supervised prototype metric has a drastic performance impairment when pruning away 40% of the data, yet the model “BC_pruning-fraction-0.4” still achieves almost the same OOD accuracy as the baseline trained on the full dataset. The core take-away is: While more analyses would be necessary to investigate whether pruning indeed consistently helps on OOD performance, it seems safe to conclude that it does not hurt OOD performance on the investigated datasets compared to an accuracy-matched baseline. (The control setting, pruning away hard examples shown in orange, leads to much lower IID accuracies and consequently also lower OOD accuracies.) For reference, the numerical results from Figure 25(a) are also shown in Table 5.

Figures 25(b), 25(c) and 25(d) focus on a related question, the question of whether models show human-like behavior on OOD datasets. Figure 25(b) shows that two models pruned using our self-supervised prototype metric somewhat more closely match human accuracies compared to the baseline in blue; Figures 25(c) and 25(d) specifically focus on image-level consistency with human responses. In terms of overall consistency (c), the baseline scores best; in terms of error consistency pruned models outperform the VISSL-trained baseline. For details on the metrics we kindly refer the interested reader to . Numerical results are again also shown in a Table (Table 6).

Finally, in Figure 26 we observe that best-case pruning leads to slightly higher shape bias as indicated by green vertical lines plotting the average shape bias across categories, which are shifted to the left of the baseline; while worst-case pruning in orange is shifted to the right. An outlier is the torchvision-trained model in purple with a very strong texture bias; we attribute this to data augmentation differences between VISSL and torchvision training.