Dataset Distillation

Tongzhou Wang, Jun-Yan Zhu, Antonio Torralba, Alexei A. Efros

Introduction

Hinton et al. (2015) proposed network distillation as a way to transfer the knowledge from an ensemble of many separately-trained networks into a single, typically compact network, performing a type of model compression. In this paper, we are considering a related but orthogonal task: rather than distilling the model, we propose to distill the dataset. Unlike network distillation, we keep the model fixed but encapsulate the knowledge of the entire training dataset, which typically contains thousands to millions of images, into a small number of synthetic training images. We show that we can go as low as one synthetic image per category, training the same model to reach surprisingly good performance on these synthetic images. For example in Figure 1a, we compress 60,00060,000 training images of MNIST digit dataset into only 1010 synthetic images (one per class), given a fixed network initialization. Training the standard LeNet (LeCun et al., 1998) on these 1010 images yields test-time MNIST recognition performance of 94%94\%, compared to 99%99\% for the original dataset. For networks with unknown random weights, 100100 synthetic images train to 80%80\% with a few gradient descent steps. We name our method Dataset Distillation and these images distilled images.

But why is dataset distillation useful? There is the purely scientific question of how much data is encoded in a given training set and how compressible it is? Moreover, given a few distilled images, we can now “load up" a given network with an entire dataset-worth of knowledge much more efficiently, compared to traditional training that often uses tens of thousands of gradient descent steps.

A key question is whether it is even possible to compress a dataset into a small set of synthetic data samples. For example, is it possible to train an image classification model on synthetic images out of the manifold of natural images? Conventional wisdom would suggest that the answer is no, as the synthetic training data may not follow the same distribution of the real test data. Yet, in this work, we show that this is indeed possible. We present a new optimization algorithm for synthesizing a small number of synthetic data samples not only capturing much of the original training data but also tailored explicitly for fast model training in only a few gradient steps. To achieve our goal, we first derive the network weights as a differentiable function of our synthetic training data. Given this connection, instead of optimizing the network weights for a particular training objective, we optimize the pixel values of our distilled images. However, this formulation requires access to the initial weights of the network. To relax this assumption, we develop a method for generating distilled images for randomly initialized networks. To further boost performance, we propose an iterative version, where we obtain a sequence of distilled images and these distilled images can be trained with multiple epochs (passes). Finally, we study a simple linear model, deriving a lower bound on the size of distilled data required to achieve the same performance as training on the full dataset.

We demonstrate that a handful of distilled images can be used to train a model with a fixed initialization to achieve surprisingly high performance. For networks pre-trained on other tasks, our method can find distilled images for fast model fine-tuning. We test our method on several initialization settings: fixed initialization, random initialization, fixed pre-trained weights, and random pre-trained weights, as well as two training objectives: image classification and malicious dataset poisoning attack. Extensive experiments on four publicly available datasets, MNIST, CIFAR10, PASCAL-VOC, and CUB-200, show that our approach often outperforms existing methods. Please check out our code and website for more details.

Related Work

Knowledge distillation. The main inspiration for this paper is network distillation (Hinton et al., 2015), a widely used technique in ensemble learning (Radosavovic et al., 2018) and model compression (Ba & Caruana, 2014; Romero et al., 2015; Howard et al., 2017). While network distillation aims to distill the knowledge of multiple networks into a single model, our goal is to compress the knowledge of an entire dataset into a few synthetic training images. Similar to our approach, data-free knowledge distillation also optimizes synthetic data samples, but with a different objective of matching activation statistics of a teacher model in knowledge distillation (Lopes et al., 2017). Our method is also related to the theoretical concept of teaching dimension, which specifies the size of dataset necessary to teach a target model to a learner (Shinohara & Miyano, 1991; Goldman & Kearns, 1995). However, methods (Zhu, 2013; 2015) inspired by this concept need the existence of target models, which our method does not require.

Dataset pruning, core-set construction, and instance selection. Another way to distill knowledge is to summarize the entire dataset by a small subset, either by only using the “valuable” data for model training (Angelova et al., 2005; Felzenszwalb et al., 2010; Lapedriza et al., 2013) or by only labeling the “valuable” data via active learning (Cohn et al., 1996; Tong & Koller, 2001). Similarly, core-set construction (Tsang et al., 2005; Har-Peled & Kushal, 2007; Bachem et al., 2017; Sener & Savarese, 2018) and instance selection (Olvera-López et al., 2010) methods aim to select a subset of the entire training data, such that models trained on the subset will perform as well as the model trained on full dataset. For example, solutions to many classical linear learning algorithms, e.g., Perceptron (Rosenblatt, 1957) and SVMs (Hearst et al., 1998), are weighted sums of a subset of training examples, which can be viewed as core-sets. However, algorithms constructing these subsets require many more training examples per category than we do, in part because their “valuable” images have to be real, whereas our distilled images are exempt from this constraint.

Gradient-based hyperparameter optimization. Our work bears similarity with gradient-based hyperparameter optimization techniques, which compute the gradient of hyperparameter w.r.t. the final validation loss by reversing the entire training procedure (Bengio, 2000; Domke, 2012; Maclaurin et al., 2015; Pedregosa, 2016). We also backpropagate errors through optimization steps. However, we use only training set data and focus more heavily on learning synthetic training data rather than tuning hyperparameters. To our knowledge, this direction has only been slightly touched on previously (Maclaurin et al., 2015). We explore it in greater depth and demonstrate the idea of dataset distillation in various settings. More crucially, our distilled images work well across random initialization weights, not possible by prior work.

Understanding datasets. Researchers have presented various approaches for understanding and visualizing learned models (Zeiler & Fergus, 2014; Zhou et al., 2015; Mahendran & Vedaldi, 2015; Bau et al., 2017; Koh & Liang, 2017). Unlike these approaches, we are interested in understanding the intrinsic properties of the training data rather than a specific trained model. Analyzing training datasets has, in the past, been mainly focused on the investigation of bias in datasets (Ponce et al., 2006; Torralba & Efros, 2011). For example, Torralba & Efros (2011) proposed to quantify the “value” of dataset samples using cross-dataset generalization. Our method offers a new perspective for understanding datasets by distilling full datasets into a few synthetic samples.

Approach

Given a model and a dataset, we aim to obtain a new, much-reduced synthetic dataset which performs almost as well as the original dataset. We first present our main optimization algorithm for training a network with a fixed initialization with one gradient descent (GD) step (Section 3.1). In Section 3.2, we derive the resolution to a more challenging case, where initial weights are random rather than fixed. In Section 3.3, we further study a linear network case to help readers understand both the property and limitation of our method. We also discuss the initial weights distribution with which our method can work well. In Section 3.4, we extend our approach to more than one gradient descent steps and more than one epoch (pass). Finally, Section 3.5 and Section 3.6 demonstrate how to obtain distilled images with different initialization distributions and learning objectives.

Standard training usually applies minibatch stochastic gradient descent or its variants. At each step tt, a minibatch of training data xt={xt,j}j=1n\mathbf{x}_{t}=\{x_{t,j}\}_{j=1}^{n} is sampled to update the current parameters as

2 Distillation for Random Initializations

Unfortunately, the above distilled data optimized for a given initialization do not generalize well to other initializations. The distilled data often look like random noise (e.g., in Figure 2(a)) as it encodes the information of both training dataset x\mathbf{x} and a particular network initialization θ0\theta_{0}. To address this issue, we turn to calculate a small number of distilled data that can work for networks with random initializations from a specific distribution. We formulate the optimization problem as follows:

where the network initialization θ0\theta_{0} is randomly sampled from a distribution p(θ0)p(\theta_{0}). During our optimization, the distilled data are optimized to work well for randomly initialized networks. Algorithm 1 illustrates our main method. In practice, we observe that the final distilled data generalize well to unseen initializations. In addition, these distilled images often look quite informative, encoding the discriminative features of each category (e.g., in Figure 3).

3 Analysis of a Simple Linear Case

4 Multiple Gradient Descent Steps and Multiple Epochs

We extend Algorithm 1 to more than one gradient descent steps by changing Line 8 to multiple sequential GD steps each on a different batch of distilled data and learning rates, i.e., each step ii is

and changing Line 11 to backpropagate through all steps. However, naively computing gradients is memory and computationally expensive. Therefore, we exploit a recent technique called back-gradient optimization, which allows for significantly faster gradient calculation of such updates in reverse-mode differentiation. Specifically, back-gradient optimization formulates the necessary second order terms into efficient Hessian-vector products (Pearlmutter, 1994), which can be easily calculated with modern automatic differentiation systems such as PyTorch (Paszke et al., 2017). For further algorithmic details, we refer readers to prior work (Domke, 2012; Maclaurin et al., 2015).

Multiple epochs. To further improve the performance, we train the network for multiple epochs (passes) over the same sequence of distilled data. In other words, for each epoch, our method cycles through all GD steps, where each step is associated with a batch of distilled data. We do not tie the trained learning rates across epochs as later epochs often use smaller learning rates. In Section 4.1, we find that using multiple steps and multiple epochs is more effective than using just one on neural networks, with the total amount of distilled data fixed.

5 Distillation with Different Initializations

Inspired by the analysis of the simple linear case in Section 3.3, we aim to focus on initial weights distributions p(θ)p(\theta) that yield similar local conditions over the data distribution. In this work, we focus on the following four practical choices:

Random initialization: Distribution over random initial weights, e.g., He Initialization (He et al., 2015) and Xavier Initialization (Glorot & Bengio, 2010) for neural networks.

Fixed initialization: A particular fixed network initialized by the method above.

Random pre-trained weights: Distribution over models pre-trained on other tasks or datasets, e.g., AlexNet (Krizhevsky et al., 2012) networks trained on ImageNet (Deng et al., 2009).

Fixed pre-trained weights: A particular fixed network pre-trained on other tasks and datasets.

Distillation with pre-trained weights. Such learned distilled data essentially fine-tune weights pre-trained on one dataset to perform well for a new dataset, thus bridging the gap between the two domains. Domain mismatch and dataset bias represent a challenging problem in machine learning today (Torralba & Efros, 2011). Extensive prior work has been proposed to adapt models to new tasks and datasets (Daume III, 2007; Saenko et al., 2010). In this work, we characterize the domain mismatch via distilled data. In Section 4.2, we show that a small number of distilled images are sufficient to quickly adapt CNN models to new datasets and tasks.

6 Distillation with Different Objectives

Distilled data learned with different learning objectives can train models to exhibit different desired behaviors. We have already mentioned image classification as one of the applications, where distilled images help to train accurate classifiers. Below, we introduce a different learning objective to further demonstrate the flexibility of our method.

Distillation for malicious data poisoning. For example, our approach can be used to construct a new form of data poisoning attack. To illustrate this idea, we consider the following scenario. When a single GD step is applied with our synthetic adversarial data, a well-behaved image classifier catastrophically forgets one category but still maintains high accuracy on other categories.

where p(θ0)p(\theta_{0}) is the distribution over random weights of well-optimized classifiers. Trained on a distribution of such classifiers, the distilled images do not require access to the exact model weights and thus can generalize to unseen models. In our experiments, the malicious distilled images are trained on 20002000 well-optimized models and evaluated on 200 held-out ones.

Compared to prior data poisoning attacks (Biggio et al., 2012; Li et al., 2016; Muñoz-González et al., 2017; Koh & Liang, 2017), our approach crucially does not require the poisoned training data to be stored and trained on repeatedly. Instead, our method attacks the model training in one iteration and with only a few data. This advantage makes our method potentially effective for online training algorithms and useful for the case where malicious users hijack the data feeding pipeline for only one gradient step (e.g., one network transmission). In Section 4.2, we show that a single batch of distilled data applied in one step can successfully attack well-optimized neural network models. This setting can be viewed as distilling the knowledge of a specific category into data.

Experiments

We report image classification results on MNIST (LeCun, 1998) and CIFAR10 (Krizhevsky & Hinton, 2009). For MNIST, the distilled images are trained with LeNet (LeCun et al., 1998), which achieves about 99%99\% test accuracy if fully trained. For CIFAR10, we use a network architecture (Krizhevsky, 2012) that achieves around 80%80\% test accuracy if fully trained. For random initializations and random pre-trained weights, we report means and standard deviations over 200200 held-out models, unless otherwise specified. The code and full results can be found at our website.

Baselines. For each experiment, in addition to baselines specific to the setting, we generally compare our method against baselines trained with data derived or selected from real training images:

Random real images: We randomly sample the same number of real images per category.

Optimized real images: We sample different sets of random real images as above, and choose the top 20%20\% best performing sets.

kk-means: We apply kk-means clustering to each category, and use the cluster centroids as training images.

Average real images: We compute the average image for each category, which is reused in different GD steps.

For these baselines, we perform each evaluation on 200200 held-out models with all combinations of learning rate∈{learned learning rate with our method,0.001,0.003,0.01,0.03,0.1,0.3}\text{learning rate}\in\{\text{learned learning rate with our method},0.001,0.003,0.01,0.03,0.1,0.3\} and #epochs∈{1,3,5}\text{\#epochs}\in\{1,3,5\}, and report results using the best performing combination. Please see the appendix Section S-1 for more details about training and baselines.

Fixed initialization. With access to initial network weights, distilled images can directly train a fixed network to reach high performance. For example, 1010 distilled images can boost the performance of a neural network with an initial accuracy 12.90%{{}{}{}{}{}}12.90\% to a final accuracy 93.76%{{}{}{}{}{}}93.76\% on MNIST (Figure 2(a)). Similarly, 100100 images can train a network with an initial accuracy 8.82%{{}{}{}{}{}}8.82\% to 54.03%{{}{}{}{}{}}54.03\% test accuracy on CIFAR10 (Figure 2(b)). This result suggests that even only a few distilled images have enough capacity to distill part of the dataset.

Random initialization. Trained with randomly sampled initializations using Xavier Initialization (Glorot & Bengio, 2010), the learned distilled images do not need to encode information tailored for a particular starting point and thus can represent meaningful content independent of network initializations. In Figure 3, we see that such distilled images reveal the discriminative features of corresponding categories: e.g., the ship image in Figure 2(d). These 100100 images can train randomly initialized networks to 36.79%{{}{}{}{}{}}36.79\% average test accuracy on CIFAR10. Similarly, for MNIST, the 100100 distilled images shown in Figure 2(c) can train randomly initialized networks to 79.50%{{}{}{}{}{}}79.50\% test accuracy.

Multiple gradient descent steps and multiple epochs. In Figure 3, distilled images are learned for 1010 GD steps applied in 33 epochs, leading to a total of 100100 images (with each step containing one image per category). Images used in early steps tend to look noisier. However, in later steps, the distilled images gradually look like real data and share the discriminative features for these categories. Figure 4(a) shows that using more steps significantly improves the results. Figure 4(b) shows a similar but slower trend as the number of epochs increases. We observe that longer training (i.e., more epochs) can help the model learn all the knowledge from the distilled images, but the performance is eventually limited by the total number of images. Alternatively, we can train the model with one GD step but a big batch size. Section 3.3 has shown theoretical limitations of using only one step in a simple linear case. In Figure 5, we observe that using multiple steps drastically outperforms the single step method, given the same number of distilled images.

Table 1 summarizes the results of our method and all baselines. Our method with both fixed and random initializations outperforms all the baselines on CIFAR10 and most of the baselines on MNIST.

2 Distillation with Different Initializations and Objectives

Next, we show two extended settings of our main algorithm discussed in Section 3.5 and Section 3.6. Both cases assume that the initial weights are random but pre-trained on the same dataset. We train the distilled images on 20002000 random pre-trained models and evaluate them on unseen models.

Fixed and random pre-trained weights on digits. As shown in Section 3.5, we can optimize distilled images to quickly fine-tune pre-trained models for a new dataset. Table 3 shows that our method is more effective than various baselines on adaptation between three digits datasets: MNIST, USPS (Hull, 1994), and SVHN (Netzer et al., 2011). We also compare our method against a state-of-the-art few-shot domain adaptation method (Motiian et al., 2017). Although our method uses the entire training set to compute the distilled images, both methods use the same number of images to distill the knowledge of target dataset. Prior work (Motiian et al., 2017) is outperformed by our method with fixed pre-trained weights on all the tasks, and by our method with random pre-trained weights on two of the three tasks. This result shows that our distilled images effectively compress the information of target datasets.

Fixed pre-trained weights on ImageNet. In Table 3, we adapt a widely used AlexNet model (Krizhevsky et al., 2012) pre-trained on ImageNet (Deng et al., 2009) to image classification on PASCAL-VOC (Everingham et al., 2010) and CUB-200 (Wah et al., 2011) datasets. Using only one distilled image per category, our method outperforms baselines significantly. Our method is on par with fine-tuning on the full datasets with thousands of images.

Random pre-trained weights and a malicious data-poisoning objective. Section 3.6 shows that our method can construct a new type of data poisoning, where an attacker can apply just one GD step with a few malicious data to manipulate a well-trained model. We train distilled images to make well-optimized neural networks to misclassify an attacked category as another target category within only one GD step. Our method requires no access to the exact weights of the model. In Figure 6(b), we evaluate our method on 200200 held-out models, against various baselines using data derived from real images and incorrect labels. For baselines, we apply one GD step using the same numbers of images with modified labels (i.e., the attacked category images are labeled as target category) and report the highest overall accuracy w.r.t. the modified labels while misclassifying ≥10%\geq 10\% attacked category as target category. This avoids results with learning rates too low to change model behavior at all. While some baselines perform similarly well as our method on MNIST, our method significantly outperforms all the baselines on CIFAR10.

Discussion

In this paper, we have presented dataset distillation for compressing the knowledge of entire training data into a few synthetic training images. We can train a network to reach high performance with a small number of distilled images and several gradient descent steps. Finally, we demonstrate two extended settings including adapting pre-trained models to new datasets and performing a malicious data-poisoning attack. In the future, we plan to extend our method to compressing large-scale visual datasets such as ImageNet and other types of data (e.g., audio and text). Also, our current method is sensitive to the distribution of initializations. We would like to investigate other initialization strategies, with which dataset distillation can work well.

Acknowledgments This work was supported in part by NSF 1524817 on Advancing Visual Recognition with Feature Visualizations, NSF IIS-1633310, and Berkeley Deep Drive.

References

S-1 Supplementary Material

In our experiments, we disable dropout layers in the networks due to the randomness and computational cost they introduce in distillation. Moreover, we initialize the distilled learning rates with a constant between 0.0010.001 and 0.020.02 depending on the task, and use the Adam solver (Kingma & Ba, 2015) with a learning rate of 0.0010.001. For random initialization and random pre-trained weights, we sample 44 to 1616 initial weights in each optimization step. We run all the experiments on NVIDIA Titan Xp and V100 GPUs. We use one GPU for fixed initial weights and four GPUs for random initial weights. Each training typically takes 11 to 44 hours.

Below we describe the details of our baselines using real training images.

Random real images: We randomly sample the same number of real images per category. We evaluate the performance over 1010 randomly sampled sets.

Optimized real images: We sample 5050 sets of real images using the procedure above, pick 1010 sets that achieve the best performance on 2020 held-out models and 10241024 randomly chosen training images, and evaluate the performance of these 1010 sets.

kk-means: For each category, we use kk-means clustering to extract the same number of cluster centroids as the number of distilled images in our method. We evaluate the method over 1010 runs.

Average real images: We compute the average image of all the images in each category, which is reused in different GD steps. We evaluate the model only once because average images are deterministic.

To enforce our optimized learning rate to be positive, we apply softplus to a scalar trained parameter.