Mean teachers are better role models: Weight-averaged consistency targets improve semi-supervised deep learning results

Antti Tarvainen, Harri Valpola

Introduction

Deep learning has seen tremendous success in areas such as image and speech recognition. In order to learn useful abstractions, deep learning models require a large number of parameters, thus making them prone to over-fitting (Figure 1a). Moreover, adding high-quality labels to training data manually is often expensive. Therefore, it is desirable to use regularization methods that exploit unlabeled data effectively to reduce over-fitting in semi-supervised learning.

When a percept is changed slightly, a human typically still considers it to be the same object. Correspondingly, a classification model should favor functions that give consistent output for similar data points. One approach for achieving this is to add noise to the input of the model. To enable the model to learn more abstract invariances, the noise may be added to intermediate representations, an insight that has motivated many regularization techniques, such as Dropout . Rather than minimizing the classification cost at the zero-dimensional data points of the input space, the regularized model minimizes the cost on a manifold around each data point, thus pushing decision boundaries away from the labeled data points (Figure 1b).

Since the classification cost is undefined for unlabeled examples, the noise regularization by itself does not aid in semi-supervised learning. To overcome this, the Γ\Gamma model evaluates each data point with and without noise, and then applies a consistency cost between the two predictions. In this case, the model assumes a dual role as a teacher and a student. As a student, it learns as before; as a teacher, it generates targets, which are then used by itself as a student for learning. Since the model itself generates targets, they may very well be incorrect. If too much weight is given to the generated targets, the cost of inconsistency outweighs that of misclassification, preventing the learning of new information. In effect, the model suffers from confirmation bias (Figure 1c), a hazard that can be mitigated by improving the quality of targets.

There are at least two ways to improve the target quality. One approach is to choose the perturbation of the representations carefully instead of barely applying additive or multiplicative noise. Another approach is to choose the teacher model carefully instead of barely replicating the student model. Concurrently to our research, Miyato et al. 2017 have taken the first approach and shown that Virtual Adversarial Training can yield impressive results. We take the second approach and will show that it too provides significant benefits. To our understanding, these two approaches are compatible, and their combination may produce even better outcomes. However, the analysis of their combined effects is outside the scope of this paper.

Our goal, then, is to form a better teacher model from the student model without additional training. As the first step, consider that the softmax output of a model does not usually provide accurate predictions outside training data. This can be partly alleviated by adding noise to the model at inference time , and consequently a noisy teacher can yield more accurate targets (Figure 1d). This approach was used in Pseudo-Ensemble Agreement and has lately been shown to work well on semi-supervised image classification . Laine & Aila 2016 named the method the Π\Pi model; we will use this name for it and their version of it as the basis of our experiments.

The Π\Pi model can be further improved by Temporal Ensembling , which maintains an exponential moving average (EMA) prediction for each of the training examples. At each training step, all the EMA predictions of the examples in that minibatch are updated based on the new predictions. Consequently, the EMA prediction of each example is formed by an ensemble of the model’s current version and those earlier versions that evaluated the same example. This ensembling improves the quality of the predictions, and using them as the teacher predictions improves results. However, since each target is updated only once per epoch, the learned information is incorporated into the training process at a slow pace. The larger the dataset, the longer the span of the updates, and in the case of on-line learning, it is unclear how Temporal Ensembling can be used at all. (One could evaluate all the targets periodically more than once per epoch, but keeping the evaluation span constant would require O(n2)O(n^{2}) evaluations per epoch where nn is the number of training examples.)

Mean Teacher

To overcome the limitations of Temporal Ensembling, we propose averaging model weights instead of predictions. Since the teacher model is an average of consecutive student models, we call this the Mean Teacher method (Figure 2). Averaging model weights over training steps tends to produce a more accurate model than using the final weights directly . We can take advantage of this during training to construct better targets. Instead of sharing the weights with the student model, the teacher model uses the EMA weights of the student model. Now it can aggregate information after every step instead of every epoch. In addition, since the weight averages improve all layer outputs, not just the top output, the target model has better intermediate representations. These aspects lead to two practical advantages over Temporal Ensembling: First, the more accurate target labels lead to a faster feedback loop between the student and the teacher models, resulting in better test accuracy. Second, the approach scales to large datasets and on-line learning.

More formally, we define the consistency cost JJ as the expected distance between the prediction of the student model (with weights θ\theta and noise η\eta) and the prediction of the teacher model (with weights θ′\theta^{\prime} and noise η′\eta^{\prime}).

The difference between the Π\Pi model, Temporal Ensembling, and Mean teacher is how the teacher predictions are generated. Whereas the Π\Pi model uses θ′=θ\theta^{\prime}=\theta, and Temporal Ensembling approximates f(x,θ′,η′)f(x,\theta^{\prime},\eta^{\prime}) with a weighted average of successive predictions, we define θt′\theta^{\prime}_{t} at training step tt as the EMA of successive θ\theta weights:

where α\alpha is a smoothing coefficient hyperparameter. An additional difference between the three algorithms is that the Π\Pi model applies training to θ′\theta^{\prime} whereas Temporal Ensembling and Mean Teacher treat it as a constant with regards to optimization.

We can approximate the consistency cost function JJ by sampling noise η,η′\eta,\eta^{\prime} at each training step with stochastic gradient descent. Following Laine & Aila 2016, we use mean squared error (MSE) as the consistency cost in most of our experiments.

Experiments

To test our hypotheses, we first replicated the Π\Pi model in TensorFlow as our baseline. We then modified the baseline model to use weight-averaged consistency targets. The model architecture is a 13-layer convolutional neural network (ConvNet) with three types of noise: random translations and horizontal flips of the input images, Gaussian noise on the input layer, and dropout applied within the network. We use mean squared error as the consistency cost and ramp up its weight from 0 to its final value during the first 80 epochs. The details of the model and the training procedure are described in Appendix B.1.

We ran experiments using the Street View House Numbers (SVHN) and CIFAR-10 benchmarks . Both datasets contain 32x32 pixel RGB images belonging to ten different classes. In SVHN, each example is a close-up of a house number, and the class represents the identity of the digit at the center of the image. In CIFAR-10, each example is a natural image belonging to a class such as horses, cats, cars and airplanes. SVHN contains of 73257 training samples and 26032 test samples. CIFAR-10 consists of 50000 training samples and 10000 test samples.

Tables 1 and 2 compare the results against recent state-of-the-art methods. All the methods in the comparison use a similar 13-layer ConvNet architecture. Mean Teacher improves test accuracy over the Π\Pi model and Temporal Ensembling on semi-supervised SVHN tasks. Mean Teacher also improves results on CIFAR-10 over our baseline Π\Pi model.

The recently published version of Virtual Adversarial Training by Miyato et al. 2017 performs even better than Mean Teacher on the 1000-label SVHN and the 4000-label CIFAR-10. As discussed in the introduction, VAT and Mean Teacher are complimentary approaches. Their combination may yield better accuracy than either of them alone, but that investigation is beyond the scope of this paper.

2 SVHN with extra unlabeled data

Above, we suggested that Mean Teacher scales well to large datasets and on-line learning. In addition, the SVHN and CIFAR-10 results indicate that it uses unlabeled examples efficiently. Therefore, we wanted to test whether we have reached the limits of our approach.

Besides the primary training data, SVHN includes also an extra dataset of 531131 examples. We picked 500 samples from the primary training as our labeled training examples. We used the rest of the primary training set together with the extra training set as unlabeled examples. We ran experiments with Mean Teacher and our baseline Π\Pi model, and used either 0, 100000 or 500000 extra examples. Table 3 shows the results.

3 Analysis of the training curves

The training curves on Figure 3 help us understand the effects of using Mean Teacher. As expected, the EMA-weighted models (blue and dark gray curves in the bottom row) give more accurate predictions than the bare student models (orange and light gray) after an initial period.

Using the EMA-weighted model as the teacher improves results in the semi-supervised settings. There appears to be a virtuous feedback cycle of the teacher (blue curve) improving the student (orange) via the consistency cost, and the student improving the teacher via exponential moving averaging. If this feedback cycle is detached, the learning is slower, and the model starts to overfit earlier (dark gray and light gray).

Mean Teacher helps when labels are scarce. When using 500 labels (middle column) Mean Teacher learns faster, and continues training after the Π\Pi model stops improving. On the other hand, in the all-labeled case (left column), Mean Teacher and the Π\Pi model behave virtually identically.

Mean Teacher uses unlabeled training data more efficiently than the Π\Pi model, as seen in the middle column. On the other hand, with 500k extra unlabeled examples (right column), Π\Pi model keeps improving for longer. Mean Teacher learns faster, and eventually converges to a better result, but the sheer amount of data appears to offset Π\Pi model’s worse predictions.

4 Ablation experiments

To assess the importance of various aspects of the model, we ran experiments on SVHN with 250 labels, varying one or a few hyperparameters at a time while keeping the others fixed.

Removal of noise (Figures 4(a) and 4(b)). In the introduction and Figure 1, we presented the hypothesis that the Π\Pi model produces better predictions by adding noise to the model on both sides. But after the addition of Mean Teacher, is noise still needed? Yes. We can see that either input augmentation or dropout is necessary for passable performance. On the other hand, input noise does not help when augmentation is in use. Dropout on the teacher side provides only a marginal benefit over just having it on the student side, at least when input augmentation is in use.

Sensitivity to EMA decay and consistency weight (Figures 4(c) and 4(d)). The essential hyperparameters of the Mean Teacher algorithm are the consistency cost weight and the EMA decay α\alpha. How sensitive is the algorithm to their values? We can see that in each case the good values span roughly an order of magnitude and outside these ranges the performance degrades quickly. Note that EMA decay α=0\alpha=0 makes the model a variation of the Π\Pi model, although somewhat inefficient one because the gradients are propagated through only the student path. Note also that in the evaluation runs we used EMA decay α=0.99\alpha=0.99 during the ramp-up phase, and α=0.999\alpha=0.999 for the rest of the training. We chose this strategy because the student improves quickly early in the training, and thus the teacher should forget the old, inaccurate, student weights quickly. Later the student improvement slows, and the teacher benefits from a longer memory.

Decoupling classification and consistency (Figure 4(e)). The consistency to teacher predictions may not necessarily be a good proxy for the classification task, especially early in the training. So far our model has strongly coupled these two tasks by using the same output for both. How would decoupling the tasks change the performance of the algorithm? To investigate, we changed the model to have two top layers and produce two outputs. We then trained one of the outputs for classification and the other for consistency. We also added a mean squared error cost between the output logits, and then varied the weight of this cost, allowing us to control the strength of the coupling. Looking at the results (reported using the EMA version of the classification output), we can see that the strongly coupled version performs well and the too loosely coupled versions do not. On the other hand, a moderate decoupling seems to have the benefit of making the consistency ramp-up redundant.

Changing from MSE to KL-divergence (Figure 4(f)) Following Laine & Aila 2016, we use mean squared error (MSE) as our consistency cost function, but KL-divergence would seem a more natural choice. Which one works better? We ran experiments with instances of a cost function family ranging from MSE (τ=0\tau=0 in the figure) to KL-divergence (τ=1\tau=1), and found out that in this setting MSE performs better than the other cost functions. See Appendix C for the details of the cost function family and for our intuition about why MSE performs so well.

5 Mean Teacher with residual networks on CIFAR-10 and ImageNet

In the experiments above, we used a traditional 13-layer convolutional architecture (ConvNet), which has the benefit of making comparisons to earlier work easy. In order to explore the effect of the model architecture, we ran experiments using a 12-block (26-layer) Residual Network (ResNet) with Shake-Shake regularization on CIFAR-10. The details of the model and the training procedure are described in Appendix B.2. As shown in Table 4, the results improve remarkably with the better network architecture.

To test whether the methods scales to more natural images, we ran experiments on Imagenet 2012 dataset using 10% of the labels. We used a 50-block (152-layer) ResNeXt architecture , and saw a clear improvement over the state of the art. As the test set is not publicly available, we measured the results using the validation set.

Related work

Noise regularization of neural networks was proposed by Sietsma & Dow 1991. More recently, several types of perturbations have been shown to regularize intermediate representations effectively in deep learning. Adversarial Training changes the input slightly to give predictions that are as different as possible from the original predictions. Dropout zeroes random dimensions of layer outputs. Dropconnect generalizes Dropout by zeroing individual weights instead of activations. Stochastic Depth drops entire layers of residual networks, and Swapout generalizes Dropout and Stochastic Depth. Shake-shake regularization duplicates residual paths and samples a linear combination of their outputs independently during forward and backward passes.

Several semi-supervised methods are based on training the model predictions to be consistent to perturbation. The Denoising Source Separation framework (DSS) uses denoising of latent variables to learn their likelihood estimate. The Γ\Gamma variant of Ladder Network implements DSS with a deep learning model for classification tasks. It produces a noisy student predictions and clean teacher predictions, and applies a denoising layer to predict teacher predictions from the student predictions. The Π\Pi model improves the Γ\Gamma model by removing the explicit denoising layer and applying noise also to the teacher predictions. Similar methods had been proposed already earlier for linear models and deep learning . Virtual Adversarial Training is similar to the Π\Pi model but uses adversarial perturbation instead of independent noise.

The idea of a teacher model training a student is related to model compression and distillation . The knowledge of a complicated model can be transferred to a simpler model by training the simpler model with the softmax outputs of the complicated model. The softmax outputs contain more information about the task than the one-hot outputs, and the requirement of representing this knowledge regularizes the simpler model. Besides its use in model compression, distillation can be used to harden trained models against adversarial attacks . The difference between distillation and consistency regularization is that distillation is performed after training whereas consistency regularization is performed on training time.

Consistency regularization can be seen as a form of label propagation . Training samples that resemble each other are more likely to belong to the same class. Label propagation takes advantage of this assumption by pushing label information from each example to examples that are near it according to some metric. Label propagation can also be applied to deep learning models . However, ordinary label propagation requires a predefined distance metric in the input space. In contrast, consistency targets employ a learned distance metric implied by the abstract representations of the model. As the model learns new features, the distance metric changes to accommodate these features. Therefore, consistency targets guide learning in two ways. On the one hand they spread the labels according to the current distance metric, and on the other hand, they aid the network learn a better distance metric.

Conclusion

Temporal Ensembling, Virtual Adversarial Training and other forms of consistency regularization have recently shown their strength in semi-supervised learning. In this paper, we propose Mean Teacher, a method that averages model weights to form a target-generating teacher model. Unlike Temporal Ensembling, Mean Teacher works with large datasets and on-line learning. Our experiments suggest that it improves the speed of learning and the classification accuracy of the trained network. In addition, it scales well to state-of-the-art architectures and large image sizes.

The success of consistency regularization depends on the quality of teacher-generated targets. If the targets can be improved, they should be. Mean Teacher and Virtual Adversarial Training represent two ways of exploiting this principle. Their combination may yield even better targets. There are probably additional methods to be uncovered that improve targets and trained models even further.

Acknowledgements

We thank Samuli Laine and Timo Aila for fruitful discussions about their work, Phil Bachman, Colin Raffel, and Thomas Robert for noticing errors in the previous versions of this paper and everyone at The Curious AI Company for their help, encouragement, and ideas.

References

Appendix

Appendix A Results without input augmentation

See table 5 for the results without input augmentation.

Appendix B Experimental setup

Source code for the experiments is available at https://github.com/CuriousAI/mean-teacher.

We replicated the Π\Pi model of Laine & Aila 2016 in TensorFlow , and added support for Mean Teacher training. We modified the model slightly to match the requirements of the experiments, as described in subsections B.1.1 and B.1.2. The difference between the original Π\Pi model described by Laine & Aila 2016 and our baseline Π\Pi model thus depends on the experiment. The difference between our baseline Π\Pi model and our Mean Teacher model is whether the teacher weights are identical to the student weights or an EMA of the student weights. In addition, the Π\Pi models (both the original and ours) backpropagate gradients to both sides of the model whereas Mean Teacher applies them only to the student side.

Table 6 describes the architecture of the convolutional network. We applied mean-only batch normalization and weight normalization on convolutional and softmax layers. We used Leaky ReLu with α=0.1\alpha=0.1 as the nonlinearity on each of the convolutional layers.

We used cross-entropy between the student softmax output and the one-hot label as the classification cost, and the mean square error between the student and teacher softmax outputs as the consistency cost. The total cost was the weighted sum of these costs, where the weight of classification cost was the expected number of labeled examples per minibatch, subject to the ramp-ups described below.

We trained the network with minibatches of size 100. We used Adam Optimizer for training with learning rate 0.0030.003 and parameters β1=0.9\beta_{1}=0.9, β2=0.999\beta_{2}=0.999, and ε=10−8\varepsilon=10^{-8}. In our baseline Π\Pi model we applied gradients through both teacher and student sides of the network. In Mean teacher model, the teacher model parameters were updated after each training step using an EMA with α=0.999\alpha=0.999. These hyperparameters were subject to the ramp-ups and ramp-downs described below.

We applied a ramp-up period of 40000 training steps at the beginning of training. The consistency cost coefficient and the learning rate were ramped up from 00 to their maximum values, using a sigmoid-shaped function e−5(1−x)2e^{-5(1-x)^{2}}, where x∈x\in.

We used different training settings in different experiments. In the CIFAR-10 experiment, we matched the settings of Laine & Aila 2016 as closely as possible. In the SVHN experiments, we diverged from Laine & Aila 2016 to accommodate for the sparsity of labeled data. Table 7 summarizes the differences between our experiments.

We normalized the input images with ZCA based on training set statistics.

For sampling minibatches, the labeled and unlabeled examples were treated equally, and thus the number of labeled examples varied from minibatch to minibatch.

We applied a ramp-down for the last 25000 training steps. The learning rate coefficient was ramped down to 00 from its maximum value. Adam β1\beta_{1} was ramped down to 0.50.5 from its maximum value. The ramp-downs were performed using sigmoid-shaped function 1−e−12.5x21-e^{-12.5x^{2}}, where x∈x\in. These ramp-downs did not improve the results, but were used to stay as close as possible to the settings of Laine & Aila 2016.

B.1.2 ConvNet on SVHN

We normalized the input images to have zero mean and unit variance.

When doing semi-supervised training, we used 1 labeled example and 99 unlabeled examples in each mini-batch. This was important to speed up training when using extra unlabeled data. After all labeled examples had been used, they were shuffled and reused. Similarly, after all unlabeled examples had been used, they were shuffled and reused.

We applied different values for Adam β2\beta_{2} and EMA decay rate during the ramp-up period and the rest of the training. Both of the values were 0.990.99 during the first 40000 steps, and 0.9990.999 afterwards. This helped the 250-label case converge reliably.

We trained the network for 180000 steps when not using extra unlabeled examples, for 400000 steps when using 100k extra unlabeled examples, and for 600000 steps when using 500k extra unlabeled examples.

B.1.3 The baseline ConvNet models

For training the supervised-only and Π\Pi model baselines we used the same hyperparameters as for training the Mean Teacher, except we stopped training earlier to prevent over-fitting. For supervised-only runs we did not include any unlabeled examples and did not apply the consistency cost.

We trained the supervised-only model on CIFAR-10 for 7500 steps when using 1000 images, for 15000 steps when using 2000 images, for 30000 steps when using 4000 images and for 150000 steps when using all images. We trained it on SVHN for 40000 steps when using 250, 500 or 1000 labels, and for 180000 steps when using all labels.

We trained the Π\Pi model on CIFAR-10 for 60000 steps when using 1000 labels, for 100000 steps when using 2000 labels, and for 180000 steps when using 4000 labels or all labels. We trained it on SVHN for 100000 steps when using 250 labels, and for 180000 steps when using 500, 1000, or all labels.

B.2 Residual network models

We implemented our residual network experiments in PyTorch https://github.com/pytorch/pytorch. We used different architectures for our CIFAR-10 and ImageNet experiments.

For CIFAR-10, we replicated the 26-2x96d Shake-Shake regularized architecture described in , and consisting of 4+4+4 residual blocks.

We trained the network on 4 GPUs using minibatches of 512 images, 124 of which were labeled. We sampled the images in the same way as described in the SVHN experiments above. We augmented the input images with 4x4 random translations (reflecting the pixels at borders when necessary) and random horizontal flips. (Note that following we used a larger translation size than on our earlier experiments.) We normalized the images to have channel-wise zero mean and unit variance over training data.

We trained the network using stochastic gradient descent with initial learning rate 0.2 and Nesterov momentum 0.9. We trained for 180 epochs (when training with 1000 labels) or 300 epochs (when training with 4000 labels), decaying the learning rate with cosine annealing so that it would have reached zero after 210 epochs (when 1000 labels) or 350 epochs (when 4000 labels). We define epoch as one pass through all the unlabeled examples – each labeled example was included many times in one such epoch.

We used a total cost function consisting of classification cost and three other costs: We used the dual output trick described in subsection 3.4 and Figure 4(e) with MSE cost between logits with coefficient 0.01. This simplified other hyperparameter choices and improved the results. We used MSE consistency cost with coefficient ramping up from 0 to 100.0 during the first 5 epochs, using the same sigmoid ramp-up shape as in the experiments above. We also used an L2 weight decay with coefficient 2e-4. We used EMA decay value 0.97 (when 1000 labels) or 0.99 (when 4000 labels).

B.2.2 ResNet on ImageNet

On our ImageNet evaluation runs, we used a 152-layer ResNeXt architecture consisting of 3+8+36+3 residual blocks, with 32 groups of 4 channels on the first block.

We trained the network on 10 GPUs using minibatches of 400 images, 200 of which were labeled. We sampled the images in the same way as described in the SVHN experiments above. Following , we randomly augmented images using a 10 degree rotation, a crop with aspect ratio between 3/4 and 4/3 resized to 224x224 pixels, a random horizontal flip and a color jitter. We then normalized images to have channel-wise zero mean and unit variance over training data.

We trained the network using stochastic gradient descent with maximum learning rate 0.25 and Nesterov momentum 0.9. We ramped up the learning rate linearly during the first two epochs from 0.1 to 0.25. We trained for 60 epochs, decaying the learning rate with cosine annealing so that it would have reached zero after 75 epochs.

We used a total cost function consisting of classification cost and three other costs: We used the dual output trick described in subsection 3.4 and Figure 4(e) with MSE cost between logits with coefficient 0.01. We used a KL-divergence consistency cost with coefficient ramping up from 0 to 10.0 during the first 5 epochs, using the same sigmoid ramp-up shape as in the experiments above. We also used an L2 weight decay with coefficient 5e-5. We used EMA decay value 0.9997.

B.3 Use of training, validation and test data

In the development phase of our work with CIFAR-10 and SVHN datasets, we separated 10% of training data into a validation set. We removed randomly most of the labels from the remaining training data, retaining an equal number of labels from each class. We used a different set of labels for each of the evaluation runs. We retained labels in the validation set to enable exploration of the results. In the final evaluation phase we used the entire training set, including the validation set but with labels removed.

On a real-world use case we would not possess a large fully-labeled validation set. However, this setup is useful in a research setting, since it enables a more thorough analysis of the results. To the best of our knowledge, this is the common practice when carrying out research on semi-supervised learning. By retaining the hyperparameters from previous work where possible we decreased the chance of over-fitting our results to validation labels.

In the ImageNet experiments we removed randomly most of the labels from the training set, retaining an equal number of labels from each class. For validation we used the given validation set without modifications. We used a different set of training labels for each of the evaluation runs and evaluated the results against the validation set.

Appendix C Varying between mean squared error and KL-divergence

As mentioned in subsection 3.4, we ran an experiment varying the consistency cost function between MSE and KL-divergence (reproduced in Figure 5). The exact consistency function we used was

τ∈(0,1]\tau\in(0,1] and NN is the number of classes. Taking the Taylor expansion we get

where the zeroth- and first-order terms vanish. Consequently,

The results in Figure 5 show that MSE performs better than KL-divergence or CτC_{\tau} with any τ\tau. We also tried other consistency cost weights with KL-divergence and did not reach the accuracy of MSE.

The exact reason why MSE performs better than KL-divergence remains unclear, but the form of CτC_{\tau} may help explain it. Modern neural network architectures tend to produce accurate but overly confident predictions . We can assume that the true labels are accurate, but we should discount the confidence of the teacher predictions. We can do that by having τ=1\tau=1 for the classification cost and τ<1\tau<1 for the consistency cost. Then pτp_{\tau} and qτq_{\tau} discount the confidence of the approximations while ZτZ_{\tau} keeps gradients large enough to provide a useful training signal. However, we did not perform experiments to validate this explanation.