Identifying Mislabeled Data using the Area Under the Margin Ranking

Geoff Pleiss, Tianyi Zhang, Ethan R. Elenberg, Kilian Q. Weinberger

Introduction

As deep networks become increasingly powerful, the potential improvement of novel architectures in many applications is inherently limited by data quality. In many real-world settings, datasets may contain samples that are “weakly-labeled” through proxy variables or web scraping [e.g. 60, 28, 35]. Human annotators, especially on crowdsourced platforms, can also be prone to making labeling mistakes. Even the most celebrated and highly-curated datasets, like MNIST and ImageNet , famously contain harmful examples. See Fig. 1 for suspicious examples detected by our proposed method—some are clearly mislabeled, others inherently ambiguous. Mislabeled training data are problematic for overparameterized deep networks, which can achieve zero training error even on randomly-assigned labels . If a bird is mislabeled as a dog, a model will learn overly specific filters—only applicable for this one image—which will result in overfitting and worse performance.

Our goal is to automatically identify and subsequently remove mislabeled samples from training datasets. Discarding these harmful data will reduce memorization and improve generalization. Perhaps more importantly, identifying mislabeled data allows practitioners to easily audit and curate their datasets. For example, a company might like to know about common labeling mistakes in order to reduce systematic error in their annotation pipeline. Large datasets may be too costly to manually inspect; therefore, an automated method should isolate mislabeled data with high precision and recall. Prior works have investigated multi-stage pipelines [e.g. 12, 20] or robust loss functions [e.g. 67, 61] for mislabeled sample identification. We instead wish to create a method that is fully “plug-and-play” with existing training methods for maximum compatibility and minimal implementation overhead.

To this end, we propose a novel method that identifies mislabeled data simply by observing a network’s training dynamics. Our method builds upon recent theoretical and empirical works that suggest that dynamics of SGD contain salient signals about noisy data and generalization. Consider an image of a bird accidentally mislabeled as a dog. Its memorization is the outcome of a delicate tension. During training, the gradient updates from the image itself encourage the network to (wrongly) predict the dog label, whereas gradient updates from other training images encourage predicting bird through generalization. The opposing updates between the (incorrect) assigned label and the (hidden) true class membership are ultimately reflected in the logits during training.

To capture this phenomenon, we introduce the Area Under the Margin (AUM) statistic, which measures the average difference between the logit values for a sample’s assigned class and its highest non-assigned class. Correctly-labeled data, which generalize from similarly-labeled examples, do not exhibit this tension and thus have a larger AUM than mislabeled data. To separate mislabeled samples from difficult but beneficial samples, we make a second contribution. We introduce an extra (artificial) class and purposefully assign a small percentage of threshold training data to this new class. All samples assigned to this new class are by definition mislabeled; therefore, we can use the AUM statistics of these points as a threshold to separate correctly-labeled data from mislabeled data.

The AUM statistic and threshold samples are trivially compatible with any classification network. Our package (pip install aum) can be used with any PyTorch classification model. Implementing this method simply requires logging the model’s logits during training. Training data whose AUM falls below the threshold can be confidently removed from the training set. On standard benchmark tasks, we improve upon the performance of existing methods simply by removing identified mislabeled samples. We are also able to clean many real-world datasets—including WebVision and Tiny ImageNet—for improved classification performance. Most surprisingly, removing 13%13\% of the CIFAR100 dataset results in a 1.2%1.2\% reduction in test-error for a ResNet-32 model.

Related Work

Learning with noisy data has been well studied [e.g. 17, 69, 42, 16]. Here we note several prior works, but please refer to for a complete review. Within the context of deep learning, researchers have proposed novel training architectures , model-based curriculum learning schemes , label correction , and robust loss functions . Recent work has developed theoretical guarantees for certain forms of regularization . Our work aims not only to improve model robustness with minimal changes to the training procedure, but also to increase training set quality.

Our work shares a similar pipeline with Brodley and Friedl , where we first identify mislabeled examples before training a classifier on a cleaned dataset. In particular, we identify mislabeled samples through a ranking metric coupled with a learned threshold (as suggested by ). There are many deep learning approaches that explicitly or implicitly identify mislabeled data. Some methods filter data through cross-validation , influence functions , or auxiliary networks . Many approaches use signals from training as a proxy for label quality . Arazo et al. and Li et al. fit the training losses of training data with two-component mixture models to separate clean from mislabeled data. We similarly use training dynamics; however, we rely on a metric (AUM) that is less prone to confusing difficult samples for mislabeled data. Moreover, we rely on threshold samples rather than a parametric mixture to separate mislabeled data. Some research examines a relaxed setting where a small set of data is assumed to be free of mislabeled examples . This paper considers the more restrictive setting where no subset of the training data can be trusted.

There has been recent interest in combining noisy-dataset learning with semi-supervised learning and data augmentation techniques. These approaches identify a small set of correctly-labeled data and then use the remaining untrusted data in conjunction with semi-supervised learning , pseudo-labeling , or MixUp augmentation . Our paper is primarily concerned with how to identify correctly-labeled data rather than how to re-use mislabeled data. For simplicity, we discard data identified as mislabeled. However, any approach to re-integrate mislabeled data should be compatible with our method.

Identifying Mislabeled Data

We assume our training dataset Dtrain={xi,yi}i=1N\mathcal{D}_{\text{train}}=\{\mathbf{x}_{i},y_{i}\}_{i=1}^{N} consists of two data types. A mislabeled sample is one where the assigned label does not match the input. For example, x\mathbf{x} might be a picture of a bird and its assigned label yy might be dog. A correctly-labeled sample has an assigned label that matches the ground-truth of the input. Some correctly-labeled examples might be “easy-to-learn” if they are common (e.g. y=\textscdogy=\textsc{dog}, x\mathbf{x} is a golden retriever catching a frisbee). Others might be “hard-to-learn” if they are rare-occurrences (e.g. y=\textscdogy=\textsc{dog}, x\mathbf{x} is an uncommon breed). In general, we assume both easy and hard correctly-labeled samples in Dtrain\mathcal{D}_{\text{train}} improve model generalization, whereas mislabeled examples hurt generalization. Our goal is to identify mislabeled data in Dtrain\mathcal{D}_{\text{train}} (i.e. samples that hurt generalization) simply by observing differences in training dynamics among samples.

How do we determine whether or not a training sample contributes to/against generalization (and therefore is likely to be correctly-labeled/mislabeled)? The neural network community has proposed numerous metrics to quantify generalization—from parameter norms [e.g. 37] to noise stability [e.g. 4] to sharpness of minima [e.g. 29]. In this paper we utilize a metric based on the margin of training samples, which is a well-established notion for numerous machine learning algorithms [e.g. 7, 53, 57, 59]. Recent theoretical and empirical analyses suggest that margin distributions are predictive of neural network generalization . We extend this line of work by using the margin of the final layer to identify poorly generalizing training data and ultimately improve classifier performance. While other metrics have been investigated [e.g. 27], margins are advantageous because 1) they are simple and efficient to compute during training; and 2) they naturally factorize across samples, making it possible to estimate the contribution of each data point to generalization. Designing novel metrics which satisfy these criteria remains an interesting direction for future work. In concurrent work, Northcutt et al. similarly investigate the margin for identifying mislabeled data.

Area Under the Margin (AUM) Ranking.

A negative margin corresponds to an incorrect prediction, while a positive margin corresponds to a confident correct prediction. A sample will have a very negative margin if gradient updates from similar samples oppose the sample’s (potentially incorrect) assigned label. We hypothesize that, at any given epoch during training, a mislabeled sample will in expectation have a smaller margin than a correctly-labeled sample. We capture this by averaging a sample’s margin measured at each training epoch—a metric we refer to as area under the margin (AUM):

where TT is the total number of training epochs. This metric is illustrated by Fig. 2, which plots the logits for various CIFAR10 training samples over the course of training a ResNet-32. Each of the 10 lines represents a logit for a particular class. The left and middle graphs display correctly-labeled dog examples—one that is “easy-to-learn” (low training loss) and one that is “hard-to-learn” (high training loss). For these two samples, the dog logit grows larger than all other logits. The green shaded region measures the AUM, which is positive and especially large for the easy-to-learn example. Conversely, the right plot displays a mislabeled dog training sample. For most of training the dog logit is much smaller than the bird logit (red line)—the image’s ground truth class—likely due to gradient updates from similar-looking correctly-labeled birds. Consequentially, the mislabeled dog has a very negative AUM, signified by the red area on the graph. This motivates AUM as a ranking: we expect mislabeled samples to have a smaller AUM values than correctly labeled samples (Fig. 3).

Separating clean/mislabeled AUMs with threshold samples.

In order to identify mislabeled data, we must determine a threshold that separates clean and mislabeled samples. Note that this threshold is dataset dependent: Fig. 3 displays a violin plot of AUM values for CIFAR10/100 with 40%40\% label noise. CIFAR10 samples have AUMs between -4 and 2, and most negative AUMs correspond to mislabeled samples. The values on CIFAR100 tend to be more extreme (between -7 and 5), and up to 40%40\% of clean samples have a negative AUM. With access to a trusted validation set, a threshold can be learned through a hyperparameter sweep. Here, we propose a more computationally efficient strategy to learn a threshold without validation data. During training we insert fake data—which we refer to as threshold samples—that mimic the training dynamics of mislabeled data. Data with similar or worse AUMs than threshold samples can be assumed to be mislabeled.

We construct threshold samples in a simple way: take a subset of training data and re-assign their label to a brand new class—i.e. a class that doesn’t really exist. In particular, assume that our training set has NN samples that belong to cc classes. We randomly select N/(c+1)N/(c+1) samples and re-assign their labels to c+1c+1 (adding an additional neuron to the network’s output layer for the fake c+1c+1 class). Choosing N/(c+1)N/(c+1) threshold examples ensures the extra class is as likely as other classes on average. Since the network can only raise the assigned c+1c+1 logit through memorization, we expect a small and likely negative margin for threshold samples, just as with mislabeled examples. Fig. 3 displays the AUMs of threshold samples (dashed gray lines), which are indeed smaller than correctly-labeled AUMs (blue histogram).

Using an extra class c+1c+1 for threshold samples has subtle but important properties. Firstly, all threshold samples are guaranteed to mimic mislabeled data. (Since threshold samples are constructed from potentially-mislabeled training data, assigning a random label in [1,c][1,c] could accidentally “correct” some mislabeled examples.) Moreover, the additional c+1c+1 classification task does not interfere with the primary classifiers. In this sense, the network and AUM computations are minimally affected by the threshold samples.

As a simple heuristic, we identify data with a lower AUM than the 99th99^{\text{th}} percentile threshold sample as mislabeled. (While the percentile value can be tuned through a hyperparameter sweep, we demonstrate in Appx. B that identification performance is robust to this hyperparameter.) Fig. 3 demonstrates the efficacy of this strategy. The 99th99^{\text{th}} percentile threshold AUM (thick gray line) cleanly separates correctly- and mislabeled samples on noisy CIFAR10/100.

Putting this all together,

we propose the following procedure for identifying mislabeled data:

Create a subset DTHR\mathcal{D}_{\text{THR}} of threshold samples:

Construct a modified training set Dtrain′\mathcal{D}_{\text{train}}^{\prime} that includes the threshold samples.

Train a network on Dtrain′\mathcal{D}_{\text{train}}^{\prime} until the first learning rate drop, measuring the AUM of all data.

Compute α\alpha: the 99th99^{\text{th}} percentile threshold sample AUM.

Identify mislabeled data using α\alpha as a threshold {(x,y)∈(Dtrain\textbackslashDTHR):AUMx,y≤α}\{(\mathbf{x},y)\in(\mathcal{D}_{\text{train}}\textbackslash\mathcal{D}_{\text{THR}}):\text{AUM}_{\mathbf{x},y}\leq\alpha\}.

By stopping training before the first learning rate drop, we prevent the network from converging and therefore memorizing difficult/mislabeled examples. In practice, this procedure only allows us to determine which samples in Dtrain\textbackslashDTHR\mathcal{D}_{\text{train}}\textbackslash\mathcal{D}_{\text{THR}} are mislabeled. We therefore repeat this procedure using a different set of threshold samples to identify the remaining mislabeled samples. In total the procedure takes roughly the same amount of computation as training a normal network: two networks are trained up until the first learning rate drop (roughly halfway through the training of most networks).

Experiments

We test the efficacy of AUM and threshold samples in two ways. First, we directly measure the precision and recall of our identification procedure on synthetic noisy datasets. Second, we train models on noisy datasets after removing the identified data. We use test-error as a proxy for identification performance—removing mislabeled samples should improve accuracy, whereas removing correctly-labeled samples should hurt accuracy. In all experiments we do not assume the presence of any trusted data for training or validation. (See Appx. A for all experimental details.)

We note that our method and many baselines can be used in conjunction with semi-supervised learning , pseudo-labeling , or MixUp to improve noisy-training performance. Given that our focus is identifying mislabeled examples, we consider these to be complimentary orthogonal approaches. Therefore, we do not use these additional training procedures in any of our experiments.

We use synthetically-mislabeled versions of CIFAR10/100 , where subsets of 45, ⁣00045,\!000 images are used for training. We also consider Tiny ImageNet, a 200-class subset of ImageNet with 95, ⁣00095,\!000 images resized to 64×6464\times 64. We corrupt these datasets following a uniform noise model (mislabeled samples are given labels uniformly at random). We compare against several methods from the existing literature. Arazo et al. fit the training losses with a mixture of two beta distributions (DY-Bootstrap BMM). Samples assigned to the high-loss distribution are considered mislabeled. The authors also propose using a mixture of two Gaussians (DY-Bootstrap GMM)—an approach also used by Li et al. . INCV iteratively filters training data through cross-validation. The remaining data are shared between two networks that inform each other about training samples with large loss (and are therefore likely mislabeled). We implement all methods on ResNet-32 models. For our method (AUM), as well as the BMM and GMM methods, we train networks for 150 epochs with no learning rate drops. We fit the BMM/GMMs to training losses from the last epoch. For INCV we use the publicly available implementation.

Fig. 4 displays the precision and recall of the identification methods at different noise levels. The most challenging settings for all methods are CIFAR100 and Tiny ImageNet with low noise. AUM tends to achieve the highest precision and recall in most noise settings, with precision and recall consistently ≥90%\geq 90\% in high-noise settings. It is worth emphasizing that the AUM model achieves this performance without any supervision or prior knowledge about the noise model.

2 Robust Training on Synthetic Noisy Datasets

To further evaluate AUM, we train ResNet-32 models on the noisy datasets after discarding the identified mislabeled samples. As a lower bound for test-error, we train a ResNet-32 following a Standard training procedure on the full dataset. As an upper bound, we train an Oracle ResNet-32 on only the correctly-labeled data. We do not perform early stopping since we do not assume the presence of a clean validation set. These methods use the standard ResNet training procedure described in . In addition, we compare against several baseline methods, including the INCV and DY-Boostrap We compare to the DY-Bootstrap variant proposed by that does not use MixUp to disentangle the performance benefits of mislabel identification and data augmentation. methods described above. Bootstrap and D2L interpolate the (potentially incorrect) training labels with the network’s predicted labels. MentorNet learns a weighting scheme with an LSTM. The MentorNet code release can only be used with pre-compiled 0.2/0.4/0.80.2/0.4/0.8 versions of CIFAR. Co-teaching identifies high-loss data with an auxiliary network; LDMI\text{L}_{\text{DMI}} uses a robust loss function; and Data Parameters assigns a learnable weighting parameter to each sample/class. We also test a Random Weighting scheme , where all samples are assigned a weight from a rectified normal distribution (re-drawn at every epoch). We run all baseline experiments with ResNet-32 models, using publicly available implementations from the methods’ authors.

Table 1 displays the test-error of these methods on corrupted versions of CIFAR10/100. We observe several trends. First, there is a large discrepancy between the Oracle and Standard models (up to 62%62\%). While most methods reduce this gap significantly, our identification scheme (AUM) achieves the lowest error in all settings. AUM recovers oracle performance on 20%/40%20\%/40\%-noisy CIFAR10, and surpasses oracle performance on 20%/40%20\%/40\%-noisy CIFAR100—simply by removing data from the training set. We hypothesize that AUM identifies mislabeled/ambiguous samples in the standard (uncorrupted) training set. On Tiny ImageNet (Table 2), we compare AUM against Data Parameters, DY-Bootstrap, and INCV (three of the most recent methods that use identification). As with CIFAR10/100, AUM matches (or surpasses) oracle performance in most noise settings.

3 Real-World Datasets

We test the performance of our method on two datasets where the label-noise is unknown. WebVision contains 2 million images scraped from Flickr and Google Image Search. It contains no human annotation (labels come from the scraping search queries), and therefore we expect many mislabeled examples. Similar to prior work [e.g. 12] we train on a subset, WebVision50, that contains the first 50 classes (≈ ⁣100, ⁣000\approx\!100,\!000 images). Clothing1M contains clothing images from 14 categories that are also scraped through search queries, as well as a set of “trusted” images annotated by humans. To match the size of WebVision and speed up experiments, we use a 100K subset of the full dataset (Clothing100K). For consistency with other datasets, we don’t use any trusted images for training or validation. We train ResNet-50 models from scratch on these datasets. We compare against the recent methods of Data Parameters Note that this method does not explicitly remove data—therefore, we only compare to its final test error. , DY-Bootstrap, and INCV.

We use AUM/threshold samples to remove mislabeled training samples and re-train on the cleaned dataset. In Table 3, we compare this approach to a (Standard) model trained on the full dataset. On WebVision50 we flag 17.8%17.8\% of the data as mislabeled (see Appx. C for examples). Removing these samples reduces error from 21.4%21.4\% to 19.8%19.8\%. Similarly, we identify 16.7%16.7\% mislabeled samples on Clothing100K for a similar error reduction. In comparison, we find that the DY-Bootstrap method tends to estimate less label noise than our method and is unable to reduce error over standard training. DY-Bootstrap mixes the assigned and predicted labels during training; therefore, we hypothesize that it is overconfident with its identifications. INCV tends to remove more data than our method. It achieves higher error on WebVision50, suggesting that it is pruning too aggressively on this dataset.

On the full Clothing1M dataset, AUM reduces error from 33.5%33.5\% (standard training) to 29.6%29.6\%. We note that AUM removes fewer data on the full 1M1M dataset than on the 100K100K subset (10.7%10.7\% versus 16.7%16.7\%). It is possibly more difficult to identify mislabeled data in larger datasets, as networks require more training to memorize large datasets. We pose this question for future work.

“Mostly-clean” datasets.

While identification methods should be robust in high-noise settings, they should also work for datasets with few mislabeled examples. To that end, we test our method on uncorrupted versions of CIFAR10, CIFAR100, and Tiny ImageNet. While some images might be ambiguous or mislabeled (due to their small size), we expect that most images are correctly-labeled. Additionally, we apply our method to the full ImageNet dataset using ResNet-50 models. In Table 3 we notice several trends. First, using DY-Bootstrap or INCV results in worse performance than a standard training procedure. INCV flags over a quarter of CIFAR100 and Tiny ImageNet as mislabeled, and therefore likely throws away too much data. This suggests that this method is susceptible to confusing hard (but beneficial) training data for mislabeled data. DY-Bootstrap tends to remove fewer examples; however, it is again likely that the bootstrap loss is overconfident. Data Parameters provides marginal improvement on ImageNet but has little effect on other datasets.

In contrast, our cleaning procedure reduces the error on most datasets. (See Appx. C for examples of removed images.) By removing high AUM data, we reduce the error on CIFAR100 from 33.0%33.0\% to 31.8%31.8\%. The amount of removed data differs among the datasets: 3%3\% on CIFAR10, 13%13\% on CIFAR100, and 24%24\% on Tiny ImageNet. Based on these results, we hypothesize that the AUM method removes few beneficial training examples, and primarily removes truly-mislabeled data. On the ImageNet dataset, only 2%2\% of samples are flagged, and removing these samples does not significantly change top-1 error (from 24.224.2 to 24.424.4). Given the rigorous annotation process of this dataset, it is not surprising that we find few mislabeled samples.

A list of mislabeled data flagged by AUM is available at http://bit.ly/aum_mislabeled.

4 Analysis and Ablation Studies

AUM is a running average of a sample’s margin. We find that this averaging is necessary for a stable and separable metric. Fig. 6 (left) displays the un-averaged margin for a clean sample and a noisy sample over the course of training (CIFAR10, 40%40\% noise). Note that the two margins occupy a similar range of values and intersect several times throughout training. Conversely, AUM’s running average improves the signal-to-noise ratio in the margin trajectories. In Fig. 6 (right), we see a consistent separation after the first few epochs.

Removing data according to the AUM ranking.

Our method removes data with a lower AUM than threshold samples. Here we study the effect of removing fewer samples (lower threshold), more samples (higher threshold), or random samples. We discard varying amounts of the CIFAR100 dataset and compare models trained on the resulting subsets. We examine 1) discarding data in order of their AUM ranking; and 2) according to a random permutation. Fig. 6 displays test-error as a function of dataset size. Discarding data at random strictly increases test error regardless of threshold. On the other hand, discarding data according to AUM ranking (red line) results in a distinct optimum. This optimum corresponds to the 99%99\% threshold sample (black dotted line), suggesting our proposed method identifies samples harmful to generalization and keeps data required for good performance.

Effect of data augmentation.

In all the above experiments, we train networks and compute AUM values using standard data augmentation (random image flips and crops). We note that our method is effective even without data augmentation. On CIFAR10 (40%40\% noise) with augmentation, AUM reduces error from 43%43\% to 12%12\%; without augmentation, it reduces error from 51%51\% to 20%20\%.

Robustness to architecture and hyperparameter choices.

We find that our method is robust to architecture choice. The AUM ranking achieves 98%98\% Spearman’s correlation across networks of various depth and architecture (see Appx. B for details). This suggests that the AUM statistic captures dataset-dependent properties rather than model-dependent properties. Moreover, our method is robust to the choice of threshold sample percentile. For example, if we choose the AUM threshold to be the 90th90^{\text{th}} percentile of threshold samples, the final test-error only differs by 4%4\% (relative)—see Appx. B.

5 Limitations.

AUM/threshold samples are able to identify mislabeled samples in many real-world datasets. Nevertheless, we can construct challenging synthetic scenarios for our method. One such setting is when mislabelings are extremely systematic: for example, all bird images are either correctly labeled or assigned the (incorrect) label dog (i.e. they are never mislabeled as any other class). To construct such a asymmetric noise setting, we assume some ordering of the classes [1,c][1,c]. With probability pp, we alter a sample’s assigned label yy from its ground-truth class y~\widetilde{y} to the adjacent class y~+1\widetilde{y}+1. In this noise setting, we expect the 99%99\% threshold will be too small. Imagine that birds are mislabeled as dogs with probability p=0.4p=0.4. On average, every image with a (ground-truth) bird will generate the bird prediction with ≈60%\approx 60\% confidence and dog with ≈40%\approx 40\% confidence. Compared to the uniform noise setting, correctly-labeled birds will have a smaller margin and mislabeled birds will have a larger margin. For threshold samples however, the model confidence will be ≈1/(c+1)\approx 1/(c+1)—representing the frequency of these samples. Threshold sample margins will therefore be smaller than mislabeled margins, resulting in low identification recall.

Our proposed method achieves high precision and recall with 20%20\% asymmetric noise. However, in Fig. 8 we see that recall struggles in the 30%30\% and 40%40\% noise settings. We would note that 40%40\% noise is a nearly maximal amount of noise under this noise model (50%50\% noise would essentially be random guessing), and that the BMM and GMM approaches have a similarly low recall. Fig. 8 displays test error after discarding data flagged by our method. For baselines, we compare against methods that—like our approach—have no prior knowledge of the noise model. This excludes Co-Teaching, which requires a noise estimate. Though our method outperforms others with 20%20\% noise, it is not competitive with the best method in the 40%40\% setting. A different threshold sample construction—one specifically designed for this noise model—might result in a better AUM threshold that makes our method more competitive. However, given that AUM achieves significant error reductions on real-world datasets (Table 3), we hypothesize that this particular synthetic high-noise setting is not too common in practice.

Discussion and Conclusion

This paper introduces the AUM statistic and the method of training with threshold samples. Together, these contributions reveal differences in training dynamics that identify noisy labels with high precision and recall. We observe performance improvements in real-world settings—both for relatively-clean and very noisy datasets. Moreover, AUM can be easily combined with other noisy-training methods, such as those using data augmentation or semi-supervised learning. For researchers, we believe that the training phenomena exploited by AUM and threshold samples are an exciting area for developing new methods and rigorous theory. There are several additional directions for future work, such as using the AUM ranking for curriculum learning [e.g. 10, 25, 52], or investigating whether AUM mitigates the double descent phoenomenon for naturally noisy datasets .

Importantly, the AUM method works with any classifier “out-of-the-box” without any changes to architecture or training procedure. We provide a simple package (pip install aum) that computes AUM for any PyTorch classifier. Running this method simply requires one additional round of model training, which can be easily baked into an existing model selection procedure. Even on relatively clean datasets (like CIFAR100), dataset cleaning with AUM can lead to improvements that are potentially as impactful as a thorough architecture/hyperparameter search. For practitioners, we hope that the “dataset cleaning” step with AUM becomes a regular part of the model development pipeline, as the relatively simple procedure has potential for substantial accuracy improvements.

Broader Impact

This paper introduces a method to identify mislabeled or harmful training examples. This has implications for two types of datasets. First, it can enable the widespread use of “weakly-labeled” data [e.g. 60, 28, 35], which are often cheap to acquire but have suffered from data quality issues. Secondly, it can be used to audit existing datasets such as ImageNet , which are widely used by both researchers and practitioners to benchmark new machine learning methods and applications.

Many research and commercial applications rely on standard datasets for pre-training, like ImageNet or large text corpora . Recent work demonstrates brittle properties of these datasets ; therefore, improving their quality could impact numerous downstream tasks. However, it is also worth noting that any automated identification procedure has the potential to create or amplify existing biases in these datasets. Auditing and curation might also have unintended consequences in sensitive applications that require security or data privacy. On common datasets like ImageNet/CIFAR, it is worthwhile to note if identification errors are prone to any particular biases.

Acknowledgments and Disclosure of Funding

We would like to thank Josh Shapiro for developing our open-source PyTorch package.

References

Appendix A Experiment Details

All experiments are implemented in PyTorch . Since we don’t assume the presence of trusted validation data, we do not perform early stopping. All test errors are recorded on the model from the final epoch of training.

All tables report the mean and standard deviation from 4 trials with different random seeds. On the larger datasets we only perform a single trial (noted by results without confidence intervals).

All models unless otherwise specified are ResNet-32 models that follow the training procedure of He et al. . We train the models for 300 epochs using 10−410^{-4} weight decay, SGD with Nesterov momentum, a learning rate of 0.1, and a batch size of 256. The learning rate is dropped by a factor of 10 at epochs 150 and 225. We apply standard data augmentation: random horizontal flips and random crops. The ResNet-32 model is designed for 32×3232\times 32 images. For Tiny ImageNet (which is 64×6464\times 64), we add a stride of 2 to the initial convolution layer.

When computing the AUM to identify mislabeled data, we train these models up until the first learning rate drop (150 epochs). We additionally drop the batch size to 64 to increase the amount of variance in SGD. We find that this variance decreases the amount of memorization, which makes the AUM metric more salient. All other hyperparameters are consistent with the original training scheme.

After removing samples identified by AUM/threshold samples, we modify the batch size so that the network keeps the same number of iterations as with the full dataset. For example, if we remove 25%25\% of the data, we would modify the batch size to be 192192 (down from 256256).

WebVision50 and Clothing100K.

We train ResNet-50 models on this dataset from scratch. Almost all training details are consistent with He et al. —10−410^{-4} weight decay, SGD with Nesterov momentum, initial learning rate of 0.10.1, and a batch size of 256256. We apply standard data augmentation: random horizontal flips, random crops, and random scaling. The only difference is the length of training. Because this dataset is smaller than ImageNet, we train the models for 180 epochs. We drop the learning rate at epochs 60 and 120 by a factor of 10.

When computing the AUM, we train models up until the first learning rate drop (60 epochs) with a batch size of 256256. As with the smaller datasets, we keep the number of training iterations constant after removing high AUM examples.

Clothing1M and ImageNet.

The ImageNet procedure exactly matches the procedure for WebVision and Clothing1M, except that we only train for 9090 epochs, with learning rate drops at 3030 and 6060. The AUM is computed up until epoch 3030.

Appendix B Additional Ablation Studies

Since AUM functions like a ranking statistic, we compute the Spearman’s correlation coefficient between the different networks. In Fig. S1 (far left) we compare CIFAR10 (40%40\% noise) AUM values computed from ResNet and DenseNet models of various depths. We find that the AUM ranking is essentially the same across these networks, with >98%>98\% correlation between all pairs of networks. It is worth noting that AUM achieves this consistency in part because it is a running average across all epochs. Without this running average, the margin of samples only achieves roughly 75%75\% correlation (middle left plot). Finally, AUM is more consistent than other metrics used to identify mislabeled samples. The training loss (middle right plot), used by Arazo et al. , achieves 75%75\% correlation. Validation loss (far right plot), used by INCV , achieves 40%40\% correlation. These metrics are more susceptible to network variance, which in part explains why AUM achieves higher identification performance.

Robustness against Threshold Sample Percentile.

As a simple heuristic, we suggest using the 99th99^{\text{th}} percentile of threshold sample AUMs to separate clean data from mislabeled data (see Sec. 3). However, we note that AUM performance is robust to this choice of hyperparameter. In Fig. S2 we plot the test-error of AUM-cleaned models that use different threshold sample percentile values. On (unmodified) CIFAR100, we note that the final test error is virtually un-impacted by this hyperparameter. With 40%40\% label noise, higher percentile values typically correspond to better performance. Nevertheless, the difference between high-vs-low percentiles is relatively limited: the 90%90\%-percentile test-error is 41%41\%, whereas the 99%99\%-percentile test-error is 39%39\%.

Appendix C More Results for Real-World Datasets

Fig. S3 displays the empirical AUM densities on the real-world datasets. Unlike the synthetic mislabeled datasets (Fig. 3, main text) these datasets do not exhibit bimodal behavior. The 99%99\% threshold sample—represented by a gray line—differs for all datasets.

Example removed images.

Fig. S4 displays high AUM images for CIFAR10, CIFAR100, and Tiny ImageNet. Fig. S5 and Fig. S6 display high AUM images for WebVision50, and Clothing1M, respectively.