There Are Many Consistent Explanations of Unlabeled Data: Why You Should Average

Ben Athiwaratkun, Marc Finzi, Pavel Izmailov, Andrew Gordon Wilson

Introduction

Recent advances in deep unsupervised learning, such as generative adversarial networks (GANs) (Goodfellow et al., 2014), have led to an explosion of interest in semi-supervised learning. Semi-supervised methods make use of both unlabeled and labeled training data to improve performance over purely supervised methods. Semi-supervised learning is particularly valuable in applications such as medical imaging, where labeled data may be scarce and expensive (Oliver et al., 2018).

Currently the best semi-supervised results are obtained by consistency-enforcing approaches (Bachman et al., 2014; Laine and Aila, 2017; Tarvainen and Valpola, 2017; Miyato et al., 2017; Park et al., 2017). These methods use unlabeled data to stabilize their predictions under input or weight perturbations. Consistency-enforcing methods can be used at scale with state-of-the-art architectures. For example, the recent Mean Teacher (Tarvainen and Valpola, 2017) model has been used with the Shake-Shake (Gastaldi, 2017) architecture and has achieved the best semi-supervised performance on the consequential CIFAR benchmarks.

This paper is about conceptually understanding and improving consistency-based semi-supervised learning methods. Our approach can be used as a guide for exploring how loss geometry interacts with training procedures in general. We provide several novel observations about the training objective and optimization trajectories of the popular Π\Pi (Laine and Aila, 2017) and Mean Teacher (Tarvainen and Valpola, 2017) consistency-based models. Inspired by these findings, we propose to improve SGD solutions via stochastic weight averaging (SWA) (Izmailov et al., 2018), a recent method that averages weights of the networks corresponding to different training epochs to obtain a single model with improved generalization. On a thorough empirical study we show that this procedure achieves the best known semi-supervised results on consequential benchmarks. In particular:

We show in Section 3.1 that a simplified Π\Pi model implicitly regularizes the norm of the Jacobian of the network outputs with respect to both its inputs and its weights, which in turn encourages flatter solutions. Both the reduced Jacobian norm and flatness of solutions have been related to generalization in the literature (Sokolić et al., 2017; Novak et al., 2018; Chaudhari et al., 2016; Schmidhuber and Hochreiter, 1997; Keskar et al., 2017; Izmailov et al., 2018). Interpolating between the weights corresponding to different epochs of training we demonstrate that the solutions of Π\Pi and Mean Teacher models are indeed flatter along these directions (Figure 1).

In Section 3.2, we compare the training trajectories of the Π\Pi, Mean Teacher, and supervised models and find that the distances between the weights corresponding to different epochs are much larger for the consistency based models. The error curves of consistency models are also wider (Figure 1), which can be explained by the flatness of the solutions discussed in section 3.1. Further we observe that the predictions of the SGD iterates can differ significantly between different iterations of SGD.

We observe that for consistency-based methods, SGD does not converge to a single point but continues to explore many solutions with high distances apart. Inspired by this observation, we propose to average the weights corresponding to SGD iterates, or ensemble the predictions of the models corresponding to these weights. Averaging weights of SGD iterates compensates for larger steps, stabilizes SGD trajectories and obtains a solution that is centered in a flat region of the loss (as a function of weights). Further, we show that the SGD iterates correspond to models with diverse predictions – using weight averaging or ensembling allows us to make use of the improved diversity and obtain a better solution compared to the SGD iterates. In Section 3.3 we demonstrate that both ensembling predictions and averaging weights of the networks corresponding to different training epochs significantly improve generalization performance and find that the improvement is much larger for the Π\Pi and Mean Teacher models compared to supervised training. We find that averaging weights provides similar or improved accuracy compared to ensembling, while offering the computational benefits and convenience of working with a single model. Thus, we focus on weight averaging for the remainder of the paper.

Motivated by our observations in Section 3 we propose to apply Stochastic Weight Averaging (SWA) (Izmailov et al., 2018) to the Π\Pi and Mean Teacher models. Based on our results in Section 3.3 we propose several modifications to SWA in Section 4. In particular, we propose fast-SWA, which (1) uses a learning rate schedule with longer cycles to increase the distance between the weights that are averaged and the diversity of the corresponding predictions; and (2) averages weights of multiple networks within each cycle (while SWA only averages weights corresponding to the lowest values of the learning rate within each cycle). In Section 5, we show that fast-SWA converges to a good solution much faster than SWA.

Applying weight averaging to the Π\Pi and Mean Teacher models we improve the best reported results on CIFAR-10 for 1k1k, 2k2k, 4k4k and 10k10k labeled examples, as well as on CIFAR-100 with 10k10k labeled examples. For example, we obtain 5.0%5.0\% error on CIFAR-10 with only 4k4k labels, improving the best result reported in the literature (Tarvainen and Valpola, 2017) by 1.3%1.3\%. We also apply weight averaging to a state-of-the-art domain adaptation technique (French et al., 2018) closely related to the Mean Teacher model and improve the best reported results on domain adaptation from CIFAR-10 to STL from 19.9%19.9\% to 16.8%16.8\% error.

We release our code at https://github.com/benathi/fastswa-semi-sup

Background

We briefly review semi-supervised learning with consistency-based models. This class of models encourages predictions to stay similar under small perturbations of inputs or network parameters. For instance, two different translations of the same image should result in similar predicted probabilities. The consistency of a model (student) can be measured against its own predictions (e.g. Π\Pi model) or predictions of a different teacher network (e.g. Mean Teacher model). In both cases we will say a student network measures consistency against a teacher network.

In the semi-supervised setting, we have access to labeled data DL={(xiL,yiL)}i=1NL\mathcal{D}_{L}=\{(x^{L}_{i},y^{L}_{i})\}_{i=1}^{N_{L}}, and unlabeled data DU={xiU}i=1NU\mathcal{D}_{U}=\{x^{U}_{i}\}_{i=1}^{N_{U}}.

Given two perturbed inputs x′,x′′x^{\prime},x^{\prime\prime} of xx and the perturbed weights wf′w_{f}^{\prime} and wg′w_{g}^{\prime}, the consistency loss penalizes the difference between the student’s predicted probablities f(x′;wf′)f(x^{\prime};w_{f}^{\prime}) and the teacher’s g(x′′;wg′)g(x^{\prime\prime};w_{g}^{\prime}). This loss is typically the Mean Squared Error or KL divergence:

The total loss used to train the model can be written as

where for classification LCEL_{\text{CE}} is the cross entropy between the model predictions and supervised training labels. The parameter λ>0\lambda>0 controls the relative importance of the consistency term in the overall loss.

ΠΠ\Pi Model

The Π\Pi model, introduced in Laine and Aila (2017) and Sajjadi et al. (2016), uses the student model ff as its own teacher. The data (input) perturbations include random translations, crops, flips and additive Gaussian noise. Binary dropout (Srivastava et al., 2014) is used for weight perturbation.

Mean Teacher Model

The Mean Teacher model (MT) proposed in Tarvainen and Valpola (2017) uses the same data and weight perturbations as the Π\Pi model; however, the teacher weights wgw^{g} are the exponential moving average (EMA) of the student weights wfw^{f}: wgk=α⋅wgk−1+(1−α)⋅wfkw_{g}^{k}=\alpha\cdot w_{g}^{k-1}+(1-\alpha)\cdot w_{f}^{k}. The decay rate α\alpha is usually set between 0.90.9 and 0.9990.999. The Mean Teacher model has the best known results on the CIFAR-10 semi-supervised learning benchmark (Tarvainen and Valpola, 2017).

Other Consistency-Based Models

Temporal Ensembling (TE) (Laine and Aila, 2017) uses an exponential moving average of the student outputs as the teacher outputs in the consistency term for training. Another approach, Virtual Adversarial Training (VAT) (Miyato et al., 2017), enforces the consistency between predictions on the original data inputs and the data perturbed in an adversarial direction x′=x+ϵradvx^{\prime}=x+\epsilon r_{\text{adv}}, where radv=arg⁡max⁡r:∥r∥=1KL[f(x,w)∥f(x+ξr,w)]r_{\text{adv}}=\arg\max_{r:\|r\|=1}\textrm{KL}[f(x,w)\|f(x+\xi r,w)].

Understanding Consistency-Enforcing Models

In Section 3.1, we study a simplified version of the Π\Pi model theoretically and show that it penalizes the norm of the Jacobian of the outputs with respect to inputs, as well as the eigenvalues of the Hessian, both of which have been related to generalization (Sokolić et al., 2017; Novak et al., 2018; Dinh et al., 2017a; Chaudhari et al., 2016). In Section 3.2 we empirically study the training trajectories of the Π\Pi and MT models and compare them to the training trajectories in supervised learning. We show that even late in training consistency-based methods make large training steps, leading to significant changes in predictions on test. In Section 3.3 we show that averaging weights or ensembling predictions of the models proposed by SGD at different training epochs can lead to substantial gains in accuracy and that these gains are much larger for Π\Pi and MT than for supervised training.

Isotropic perturbations investigated in this simplified Π\Pi model will not in general lie along the data manifold, and it would be more pertinent to enforce consistency to perturbations sampled from the space of natural images. In fact, we can interpret consistency with respect to standard data augmentations (which are used in practice) as penalizing the manifold Jacobian norm in the same manner as above. See Section A.5 for more details.

Penalization of the Hessian’s eigenvalues.

2 Analysis of Solutions along SGD Trajectories

In the previous section we have seen that in a simplified Π\Pi model, the consistency loss encourages lower input-output Jacobian norm and Hessian’s eigenvalues, which are related to better generalization. In this section we analyze the properties of minimizing the consistency loss in a practical setting. Specifically, we explore the trajectories followed by SGD for the consistency-based models and compare them to the trajectories in supervised training.

We train our models on CIFAR-10 using 4k4k labeled data for 180180 epochs. The Π\Pi and Mean Teacher models use 46k46k data points as unlabeled data (see Sections A.8 and A.9 for details). First, in Figure 1 we visualize the evolution of norms of the gradients of the cross-entropy term ∥∇LCE∥\|\nabla L_{\text{CE}}\| and consistency term ∥∇Lcons∥\|\nabla L_{\text{cons}}\| along the trajectories of the Π\Pi, MT, and standard supervised models (using CE loss only). We observe that ∥∇LCons∥\|\nabla L_{\text{Cons}}\| remains high until the end of training and dominates the gradient ∥∇LCE∥\|\nabla L_{\text{CE}}\| of the cross-entropy term for the Π\Pi and MT models. Further, for both the Π\Pi and MT models, ∥∇LCons∥\|\nabla L_{\text{Cons}}\| is much larger than in supervised training implying that the Π\Pi and MT models are making substantially larger steps until the end of training. These larger steps suggest that rather than converging to a single minimizer, SGD continues to actively explore a large set of solutions when applied to consistency-based methods.

For further understand this observation, we analyze the behavior of train and test errors in the region of weight space around the solutions of the Π\Pi and Mean Teacher models. First, we consider the one-dimensional rays ϕ(t)=t⋅w180+(1−t)w170, t≥0,\phi(t)=t\cdot w_{180}+(1-t)w_{170},~{}t\geq 0, connecting the weight vectors w170w_{170} and w180w_{180} corresponding to epochs 170170 and 180180 of training. We visualize the train and test errors (measured on the labeled data) as functions of the distance from the weights w170w_{170} in Figure 1. We observe that the distance between the weight vectors w170w_{170} and w180w_{180} is much larger for the semi-supervised methods compared to supervised training, which is consistent with our observation that the gradient norms are larger which implies larger steps during optimization in the Π\Pi and MT models. Further, we observe that the train and test error surfaces are much wider along the directions connecting w170w_{170} and w180w_{180} for the consistency-based methods compared to supervised training. One possible explanation for the increased width is the effect of the consistency loss on the Jacobian of the network and the eigenvalues of the Hessian of the loss discussed in Section 3.1. We also observe that the test errors of interpolated weights can be lower than errors of the two SGD solutions between which we interpolate. This error reduction is larger in the consistency models (Figure 1).

We also analyze the error surfaces along random and adversarial rays starting at the SGD solution w180w_{180} for each model. For the random rays we sample 55 random vectors dd from the unit sphere and calculate the average train and test errors of the network with weights wt1+sdw_{t_{1}}+sd for s∈s\in. With adversarial rays we evaluate the error along the directions of the fastest ascent of test or train loss dadv=∇LCE∣∣∇LCE∣∣d_{adv}=\frac{\nabla L_{CE}}{||\nabla L_{CE}||}. We observe that while the solutions of the Π\Pi and MT models are much wider than supervised training solutions along the SGD-SGD directions (Figure 1), their widths along random and adversarial rays are comparable (Figure 1, 1)

We analyze the error along SGD-SGD rays for two reasons. Firstly, in fast-SWA we are averaging solutions traversed by SGD, so the rays connecting SGD iterates serve as a proxy for the space we average over. Secondly, we are interested in evaluating the width of the solutions that we explore during training which we expect will be improved by the consistency training, as discussed in Section 3.1 and A.6. We expect width along random rays to be less meaningful because there are many directions in the parameter space that do not change the network outputs (Dinh et al., 2017b; Gur-Ari et al., 2018; Sagun et al., 2017). However, by evaluating SGD-SGD rays, we can expect that these directions corresponds to meaningful changes to our model because individual SGD updates correspond to directions that change the predictions on the training set. Furthermore, we observe that different SGD iterates produce significantly different predictions on the test data.

Neural networks in general are known to be resilient to noise, explaining why both MT, Π\Pi and supervised models are flat along random directions (Arora et al., 2018). At the same time neural networks are susceptible to targeted perturbations (such as adversarial attacks). We hypothesize that we do not observe improved flatness for semi-supervised methods along adversarial rays because we do not choose our input or weight perturbations adversarially, but rather they are sampled from a predefined set of transformations.

3 Ensembling and Weight Averaging

In Section 3.2, we observed that the Π\Pi and MT models continue taking large steps in the weight space at the end of training. Not only are the distances between weights larger, we observe these models to have higher diversity. In this setting, using the last SGD iterate to perform prediction is not ideal since many solutions explored by SGD are equally accurate but produce different predictions.

In Section 3.2 we showed that the diversity in predictions is significantly larger for the Π\Pi and Mean Teacher models compared to purely supervised learning. The diversity of these iterates suggests that we can achieve greater benefits from ensembling. We use the same CNN architecture and hyper-parameters as in Section 3.2 but extend the training time by doing 55 learning rate cycles of 3030 epochs after the normal training ends at epoch 180180 (see A.8 and A.9 for details). We sample random pairs of weights w1w_{1}, w2w_{2} from epochs 180,183,…,330180,183,\ldots,330 and measure the error reduction from ensembling these pairs of models, Cens≡12Err(w1)+12Err(w2)−Err(Ensemble(w1,w2))C_{\text{ens}}\equiv\frac{1}{2}\textrm{Err}(w_{1})+\frac{1}{2}\textrm{Err}(w_{2})-\textrm{Err}\left(\text{Ensemble}(w_{1},w_{2})\right). In Figure 2 we visualize CensC_{\text{ens}}, against the diversity of the corresponding pair of models. We observe a strong correlation between the diversity in predictions of the constituent models and ensemble performance, and therefore CensC_{\text{ens}} is substantially larger for Π\Pi and Mean Teacher models. As shown in Izmailov et al. (2018), ensembling can be well approximated by weight averaging if the weights are close by.

Weight Averaging.

First, we experiment on averaging random pairs of weights at the end of training and analyze the performance with respect to the weight distances. Using the the same pairs from above, we evaluate the performance of the model formed by averaging the pairs of weights, Cavg(w1,w2)≡12Err(w1)+12Err(w2)−Err(12w1+12w2)C_{avg}(w_{1},w_{2})\equiv\frac{1}{2}\textrm{Err}(w_{1})+\frac{1}{2}\textrm{Err}(w_{2})-\textrm{Err}\left(\frac{1}{2}w_{1}+\frac{1}{2}w_{2}\right). Note that CavgC_{avg} is a proxy for convexity: if Cavg(w1,w2)≥0C_{avg}(w_{1},w_{2})\geq 0 for any pair of points w1w_{1}, w2w_{2}, then by Jensen’s inequality the error function is convex (see the left panel of Figure 2). While the error surfaces for neural networks are known to be highly non-convex, they may be approximately convex in the region traversed by SGD late into training (Goodfellow et al., 2015). In fact, in Figure 2, we find that the error surface of the SGD trajectory is approximately convex due to Cavg(w1,w2)C_{avg}(w_{1},w_{2}) being mostly positive. Here we also observe that the distances between pairs of weights are much larger for the Π\Pi and MT models than for the supervised training; and as a result, weight averaging achieves a larger gain for these models.

In Section 3.2 we observed that for the Π\Pi and Mean Teacher models SGD traverses a large flat region of the weight space late in training. Being very high-dimensional, this set has most of its volume concentrated near its boundary. Thus, we find SGD iterates at the periphery of this flat region (see Figure 2). We can also explain this behavior via the argument of (Mandt et al., 2017). Under certain assumptions SGD iterates can be thought of as samples from a Gaussian distribution centered at the minimum of the loss, and samples from high-dimensional Gaussians are known to be concentrated on the surface of an ellipse and never be close to the mean. Averaging the SGD iterates (shown in red in Figure 2) we can move towards the center (shown in blue) of the flat region, stabilizing the SGD trajectory and improving the width of the resulting solution, and consequently improving generalization.

We observe that the improvement CavgC_{\text{avg}} from weight averaging (1.2±0.2%1.2\pm 0.2\% over MT and Π\Pi pairs) is on par or larger than the benefit CensC_{\text{ens}} of prediction ensembling (0.9±0.2%0.9\pm 0.2\%) The smaller gain from ensembling might be due to the dependency of the ensembled solutions, since they are from the same SGD run as opposed to independent restarts as in typical ensembling settings. For the rest of the paper, we focus attention on weight averaging because of its lower costs at test time and slightly higher performance compared to ensembling.

SWA and fast-SWA

In Section 3 we analyzed the training trajectories of the Π\Pi, MT, and supervised models. We observed that the Π\Pi and MT models continue to actively explore the set of plausible solutions, producing diverse predictions on the test set even in the late stages of training. Further, in section 3.3 we have seen that averaging weights leads to significant gains in performance for the Π\Pi and MT models. In particular these gains are much larger than in supervised setting.

Stochastic Weight Averaging (SWA) (Izmailov et al., 2018) is a recent approach that is based on averaging weights traversed by SGD with a modified learning rate schedule. In Section 3 we analyzed averaging pairs of weights corresponding to different epochs of training and showed that it improves the test accuracy. Averaging multiple weights reinforces this effect, and SWA was shown to significantly improve generalization performance in supervised learning. Based on our results in section 3.3, we can expect even larger improvements in generalization when applying SWA to the Π\Pi and MT models.

Notice that most of the models included in the fast-SWA average (shown in red in Figure 3, left) have higher errors than those included in the SWA average (shown in green in Figure 3, right) since they are obtained when the learning rate is high. It is our contention that including more models in the fast-SWA weight average can more than compensate for the larger errors of the individual models. Indeed, our experiments in Section 5 show that fast-SWA converges substantially faster than SWA and has lower performance variance. We analyze this result theoretically in Section A.7).

Experiments

We evaluate the Π\Pi and MT models (Section 4) on CIFAR-10 and CIFAR-100 with varying numbers of labeled examples. We show that fast-SWA and SWA improve the performance of the Π\Pi and MT models, as we expect from our observations in Section 3. In fact, in many cases fast-SWA improves on the best results reported in the semi-supervised literature. We also demonstrate that the preposed fast-SWA obtains high performance much faster than SWA. We also evaluate SWA applied to a consistency-based domain adaptation model (French et al., 2018), closely related to the MT model, for adapting CIFAR-10 to STL. We improve the best reported test error rate for this task from 19.9%19.9\% to 16.8%16.8\%.

We discuss the experimental setup in Section 5.1. We provide the results for CIFAR-10 and CIFAR-100 datasets in Section 5.2 and 5.3. We summarize our results in comparison to the best previous results in Section 5.4. We show several additional results and detailed comparisons in Appendix A.2. We provide analysis of train and test error surfaces of fast-SWA solutions along the directions connecting fast-SWA and SGD in Section A.1.

2 CIFAR-10

50k+ and 50k+∗50k+^{*} correspond to 50k50k+500k500k and 50k50k+237k∗237k^{*} settings (c) CIFAR-10 with ResNet + Shake-Shake using the short schedule (d) CIFAR-10 with ResNet + Shake-Shake using the long schedule. We evaluate the proposed fast-SWA method using the Π\Pi and MT models on the CIFAR-10 dataset (Krizhevsky, ). We use 50k50k images for training with 1k1k, 2k2k, 4k4k, 10k10k and 50k50k labels and report the top-1 errors on the test set (10k10k images). We visualize the results for the CNN and Shake-Shake architectures in Figures 4, 4, and 4. For all quantities of labeled data, fast-SWA substantially improves test accuracy in both architectures. Additionally, in Tables 2, 4 of the Appendix we provide a thorough comparison of different averaging strategies as well as results for VAT (Miyato et al., 2017), TE (Laine and Aila, 2016), and other baselines.

Note that we applied fast-SWA for VAT as well which is another popular approach for semi-supervised learning. We found that the improvement on VAT is not drastic – our base implementation obtains 11.26%11.26\% error where fast-SWA reduces it to 10.97%10.97\% (see Table 2 in Section A.2). It is possible that the solutions explored by VAT are not as diverse as in Π\Pi and MT models due to VAT loss function. Throughout the experiments, we focus on the Π\Pi and MT models as they have been shown to scale to powerful networks such as Shake-Shake and obtained previous state-of-the-art performance.

We also find that the performance gains of fast-SWA over base models are higher for the Π\Pi model compared to the MT model, which is consistent with the convexity observation in Section 3.3 and Figure 2. In the previous evaluations (see e.g. Oliver et al., 2018; Tarvainen and Valpola, 2017), the Π\Pi model was shown to be inferior to the MT model. However, with weight averaging, fast-SWA reduces the gap between Π\Pi and MT performance. Surprisingly, we find that the Π\Pi model can outperform MT after applying fast-SWA with moderate to large numbers of labeled points. In particular, the Π\Pi+fast-SWA model outperforms MT+fast-SWA on CIFAR-10 with 4k4k, 10k10k, and 50k50k labeled data points for the Shake-Shake architecture.

3 CIFAR-100 and Extra Unlabeled Data

We evaluate the Π\Pi and MT models with fast-SWA on CIFAR-100. We train our models using 5000050000 images with 10k10k and 50k50k labels using the 1313-layer CNN. We also analyze the effect of using the Tiny Images dataset (Torralba et al., 2008) as an additional source of unlabeled data.

The Tiny Images dataset consists of 8080 million images, mostly unlabeled, and contains CIFAR-100 as a subset. Following Laine and Aila (2016), we use two settings of unlabeled data, 50k50k+500k500k and 50k50k+237k∗237k^{*} where the 50k50k images corresponds to CIFAR-100 images from the training set and the +500k+500k or +237k∗+237k^{*} images corresponds to additional 500k500k or 237k237k images from the Tiny Images dataset. For the 237k∗237k^{*} setting, we select only the images that belong to the classes in CIFAR-100, corresponding to 237203237203 images. For the 500k500k setting, we use a random set of 500k500k images whose classes can be different from CIFAR-100. We visualize the results in Figure 4, where we again observe that fast-SWA substantially improves performance for every configuration of the number of labeled and unlabeled data. In Figure 5 (middle, right) we show the errors of MT, SWA and fast-SWA as a function of iteration on CIFAR-100 for the 10k10k and 50k50k+500k500k label settings. Similar to the CIFAR-10 experiments, we observe that fast-SWA reduces the errors substantially faster than SWA. We provide detailed experimental results in Table 3 of the Appendix and include preliminary results using the Shake-Shake architecture in Table 5, Section A.2.

4 Advancing State-of-the-Art

We have shown that fast-SWA can significantly improve the performance of both the Π\Pi and MT models. We provide a summary comparing our results with the previous best results in the literature in Table 1, using the 1313-layer CNN and the Shake-Shake architecture that had been applied previously. We also provide detailed results the Appendix A.2.

5 Preliminary Results on Domain Adaptation

Domain adaptation problems involve learning using a source domain XsX_{s} equipped with labels YsY_{s} and performing classification on the target domain XtX_{t} while having no access to the target labels at training time. A recent model by French et al. (2018) applies the consistency enforcing principle for domain adaptation and achieves state-of-the-art results on many datasets. Applying fast-SWA to this model on domain adaptation from CIFAR-10 to STL we were able to improve the best results reported in the literature from 19.9%19.9\% to 16.8%16.8\%. See Section A.10 for more details on the domain adaptation experiments.

Discussion

Semi-supervised learning is crucial for reducing the dependency of deep learning on large labeled datasets. Recently, there have been great advances in semi-supervised learning, with consistency regularization models achieving the best known results. By analyzing solutions along the training trajectories for two of the most successful models in this class, the Π\Pi and Mean Teacher models, we have seen that rather than converging to a single solution SGD continues to explore a diverse set of plausible solutions late into training. As a result, we can expect that averaging predictions or weights will lead to much larger gains in performance than for supervised training. Indeed, applying a variant of the recently proposed stochastic weight averaging (SWA) we advance the best known semi-supervised results on classification benchmarks.

While not the focus of our paper, we have also shown that weight averaging has great promise in domain adaptation (French et al., 2018). We believe that application-specific analysis of the geometric properties of the training objective and optimization trajectories will further improve results over a wide range of application specific areas, including reinforcement learning with sparse rewards, generative adversarial networks (Yazıcı et al., 2018), or semi-supervised natural language processing.

References

Appendix A Appendix

In this section we provide several additional plots visualizing the train and test error along different types of rays in the weight space. The left panel of Figure 6 shows how the behavior of test error changes as we add more unlabeled data points for the Π\Pi model. We observe that the test accuracy improves monotonically, but also the solutions become narrower along random rays.

The middle panel of Figure 6 visualizes the train and test error behavior along the directions connecting the fast-SWA solution (shown with squares) to one of the SGD iterates used to compute the average (shown with circles) for Π\Pi, MT and supervised training. Similarly to Izmailov et al. (2018) we observe that for all three methods fast-SWA finds a centered solution, while the SGD solution lies near the boundary of a wide flat region. Agreeing with our results in section 3.2 we observe that for Π\Pi and Mean Teacher models the train and test error surfaces are much wider along the directions connecting the fast-SWA and SGD solutions than for supervised training.

In the right panel of Figure 6 we show the behavior of train and test error surfaces along random rays, adversarial rays and directions connecting the SGD solutions from epochs 170170 and 180180 for the Mean Teacher model (see section 3.2).

In the left panel of Figure 7 we show the evolution of the trace of the gradient of the covariance of the loss

for the Π\Pi, MT and supevised training. We observe that the variance of the gradient is much larger for the Π\Pi and Mean Teacher models compared to supervised training.

In the middle and right panels of figure 7 we provide scatter plots of the improvement CC obtained from averaging weights against diversity and diversity against distance. We observe that diversity is highly correlated with the improvement CC coming from weight averaging. The correlation between distance and diversity is less prominent.

A.2 Detailed Results

In this section we report detailed results for the Π\Pi and Mean Teacher models and various baselines on CIFAR-10 and CIFAR-100 using the 1313-layer CNN and Shake-Shake.

The results using the 13-layer CNN are summarized in Tables 2 and 3 for CIFAR-10 and CIFAR-100 respectively. Tables 4 and 5 summarize the results using Shake-Shake on CIFAR-10 and CIFAR-100. In the tables Π\Pi EMA is the same method as Π\Pi, where instead of SWA we apply Exponential Moving Averaging (EMA) for the student weights. We show that simply performing EMA for the student network in the Π\Pi model without using it as a teacher (as in MT) typically results in a small improvement in the test error.

Figures 8 and 9 show the performance of the Π\Pi and Mean Teacher models as a function of the training epoch for CIFAR-10 and CIFAR-100 respectively for SWA and fast-SWA.

A.3 Effect of Learning Rate Schedules

The only hyperparameter for the fast-SWA setting is the cycle length cc. We demonstrate in Figure 10 that fast-SWA’s performance is not sensitive to cc over a wide range of cc values. We also demonstrate the performance for constant learning schedule. fast-SWA with cyclical learning rates generally converges faster due to higher variety in the collected weights.

A.4 EMA versus SWA as a Teacher

The MT model uses an exponential moving average (EMA) of the student weights as a teacher in the consistency regularization term. We consider two potential effects of using EMA as a teacher: first, averaging weights improves performance of the teacher for the reasons discussed in Sections 3.2, 3.3; second, having a better teacher model leads to better student performance which in turn further improves the teacher. In this section we try to separate these two effects. We apply EMA to the Π\Pi model in the same way in which we apply fast-SWA instead of using EMA as a teacher and compare the resulting performance to the Mean Teacher. Figure 11 shows the improvement in error-rate obtained by applying EMA to the Π\Pi model in different label settings. As we can see while EMA improves the results over the baseline Π\Pi model, the performance of Π\Pi-EMA is still inferior to that of the Mean Teacher method, especially when the labeled data is scarce. This observation suggests that the improvement of the Mean Teacher over the Π\Pi model can not be simply attributed to EMA improving the student performance and we should take the second effect discussed above into account.

Like SWA, EMA is a way to average weights of the networks, but it puts more emphasis on very recent models compared to SWA. Early in training when the student model changes rapidly EMA significantly improves performance and helps a lot when used as a teacher. However once the student model converges to the vicinity of the optimum, EMA offers little gain. In this regime SWA is a much better way to average weights. We show the performance of SWA applied to Π\Pi model in Figure 11 (left).

Since SWA performs better than EMA, we also experiment with using SWA as a teacher instead of EMA. We start with the usual MT model pretrained until epoch 150150. Then we switch to using SWA as a teacher at epoch 150150. In Figure 11 (right), our results suggest that using SWA as a teacher performs on par with using EMA as a teacher. We conjecture that once we are at a convex region of test error close to the optimum (epoch 150150), having a better teacher doesn’t lead to substantially improved performance. It is possible to start using SWA as a teacher earlier in training; however, during early epochs where the model undergoes rapid improvement EMA is more sensible than SWA as we discussed above.

A.5 Consistency Loss Approximates Jacobian Norm

In the simplified Π\Pi model with small additive data perturbations that are normally distributed, z∼N(0,I)z\sim\mathcal{N}(0,I),

We can now recognize this term as a one sample stochastic trace estimator for tr(J(xi)TJ(xi))\text{tr}(J(x_{i})^{T}J(x_{i})) with a Gaussian probe variable ziz_{i}; see Avron and Toledo (2011) for derivations and guarantees on stochastic trace estimators.

In general if we have mm samples of xx and nn sampled perturbations for each xx, then for a symmetric matrix AA with zik∼iidN(0,I)z_{ik}\stackrel{{\scriptstyle iid}}{{\sim}}N(0,I) and independent xi∼iidp(x)x_{i}\stackrel{{\scriptstyle iid}}{{\sim}}p(x),

whereas this does not hold for the opposite ordering of the sum.

Non-isotropic perturbations along data manifold

We view the standard data augmentations such as random translation (that are applied in the Π\Pi and MT models) as approximating samples of nearby elements of the data manifold and therefore differences x′−xx^{\prime}-x approximate elements of its tangent space.

A.7 Including High Learning Rate Iterates Into SWA

A.8 Network Architectures

In the experiments we use two DNN architectures – 1313 layer CNN and Shake-Shake. The architecture of 1313-layer CNN is described in Table 6. It closely follows the architecture used in (Laine and Aila, 2017; Miyato et al., 2017; Tarvainen and Valpola, 2017). We re-implement it in PyTorch and removed the Gaussian input noise, since we found having no such noise improves generalization. For Shake-Shake we use 26-2x96d Shake-Shake regularized architecture of Gastaldi (2017) with 1212 residual blocks.

A.9 Hyperparameters

In all experiments we use stochastic gradient descent optimizer with Nesterov momentum (Loshchilov and Hutter, 2016). In fast-SWA we average every the weights of the models corresponding to every third epoch. In the Π\Pi model, we back-propagate the gradients through the student side only (as opposed to both sides in (Laine and Aila, 2016)). For Mean Teacher we use α=0.97\alpha=0.97 decay rate in the Exponential Moving Average (EMA) of the student’s weights. For all other hyper-parameters we reuse the values from Tarvainen and Valpola (2017) unless mentioned otherwise.

Like in Tarvainen and Valpola (2017), we use ∥⋅∥2\|\cdot\|^{2} for divergence in the consistency loss. Similarly, we ramp up the consistency cost λ\lambda over the first 55 epochs from up to it’s maximum value of 100100 as done in Tarvainen and Valpola (2017). We use cosine annealing learning rates with no learning rate ramp up, unlike in the original MT implementation (Tarvainen and Valpola, 2017). Note that this is similar to the same hyperparameter settings as in Tarvainen and Valpola (2017) for ResNetWe use the public Pytorch code https://github.com/CuriousAI/mean-teacher as our base model for the MT model and modified it for the Π\Pi model.. We note that we use the exact same hyperparameters for the Π\Pi and MT models in each experiment setting. In contrast to the original implementation in Tarvainen and Valpola (2017) of CNN experiments, we use SGD instead of Adam.

We use the 1313-layer CNN with the short learning rate schedule. We use a total batch size of 100100 for CNN experiments with a labeled batch size of 5050 for the Π\Pi and Mean Teacher models. We use the maximum learning rate η0=0.1\eta_{0}=0.1. For Section 3.2 we run SGD only for 180180 epochs, so learning rate cycles are done. For Section 3.3 we additionally run 55 learning rate cycles and sample pairs of SGD iterates from epochs 180180-330330 corresponding to these cycles.

CIFAR-10 CNN Experiments

We use a total batch size of 100100 for CNN experiments with a labeled batch size of 5050. We use the maximum learning rate η0=0.1\eta_{0}=0.1.

CIFAR-10 ResNet + Shake-Shake

We use a total batch size of 128128 for ResNet experiments with a labeled batch size of 3131. We use the maximum learning rate η0=0.05\eta_{0}=0.05 for CIFAR-10. This applies for both the short and long schedules.

CIFAR-100 CNN Experiments

We use a total batch size of 128128 with a labeled batch size of 3131 for 10k10k and 50k50k label settings. For the settings 50k50k+500k500k and 50k50k+237k∗237k^{*}, we use a labeled batch size of 6464. We also limit the number of unlabeled images used in each epoch to 100k100k images. We use the maximum learning rate η0=0.1\eta_{0}=0.1.

CIFAR-100 ResNet + Shake-Shake

We use a total batch size of 128128 for ResNet experiments with a labeled batch size of 3131 in all label settings. For the settings 50k50k+500k500k and 50k50k+237k∗237k^{*}, we also limit the number of unlabeled images used in each epoch to 100k100k images. We use the maximum learning rate η0=0.1\eta_{0}=0.1. This applies for both the short and long schedules.

A.10 Domain Adaptation

We apply fast-SWA to the best experiment setting MT+CT+TFA for CIFAR-10 to STL according to French et al. (2018). This setting involves using confidence thresholding (CT) and also an augmentation scheme with translation, flipping, and affine transformation (TFA).

We observe that averaging every iteration converges much faster (600600 epochs instead of 30003000) and results in better test accuracy. In our experiments with semi-supervised learning averaging more often than once per epoch didn’t improve convergence or final results. We hypothesize that the improvement from more frequent averaging is a result of specific geometry of the loss surfaces and training trajectories in domain adaptation. We leave further analysis of applying fast-SWA to domain adaptation for future work.