Smooth Neighbors on Teacher Graphs for Semi-supervised Learning

Yucen Luo, Jun Zhu, Mengxi Li, Yong Ren, Bo Zhang

Introduction

As collecting a fully labeled dataset is often expensive and time-consuming, semi-supervised learning (SSL) has been extensively studied in computer vision to improve generalization performance of the classifier by leveraging limited labeled data and a large amount of unlabeled data . The success of SSL relies on the key smoothness assumption, i.e., data points close to each other are likely to have the same label. It has a special case named cluster or low density separation assumption, which states that the decision boundary should lie in low density regions, not crossing high density regions . Based on these assumptions, many traditional methods have been developed .

Recently due to the great advances of deep learning , remarkable results have been achieved on SSL . Among these works, perturbation-based methods have demonstrated great promise. Adding noise to the deep model is important to reduce overfitting and learn more robust abstractions, e.g., dropout and randomized data augmentation . In SSL, perturbation regularization aids by exploring the smoothness assumption. For example, the Manifold Tangent Classifier (MTC) trains contrastive auto-encoders to learn the data manifold and regularizes the predictions to be insensitive to local perturbations along the low-dimensional manifold. Pseudo-Ensemble and Γ\Gamma model in Ladder Network evaluate the classifiers with and without perturbations, which act as a “teacher” and a “student”, respectively. The student needs to predict consistently with the targets generated by the teacher on unlabeled data. Following the same principle, temporal ensembling, mean teacher and virtual adversarial training improve the target quality in different ways to form better teachers. All these approaches aim to fuse the inputs into coherent clusters by adding noise and smoothing the mapping function locally .

However, these methods only consider the perturbations around each single data point, while ignoring the connections between data points, therefore not fully utilizing the information in the unlabeled data structure, such as clusters or manifolds. An extreme situation may happen where the function is smooth in the vicinity of each unlabeled point but not smooth in the vacancy among them. This artifact could be avoided if the unlabeled data structure is taken into consideration. It is known that data points similar to each other (e.g., in the same class) tend to form clusters (cluster assumption). Therefore, the connections between similar data points help the fusing of clusters become tighter and more effective (see Fig. 5 for the visualization of real data).

Motivated by that, we propose Smooth Neighbors on Teacher Graphs (SNTG) that considers the connections between data points to induce smoothness on the data manifold. By learning a teacher graph based on the targets generated by the teacher, our model encourages invariance when some perturbations are added to the neighboring points on the graph. Since deep networks have a hierarchical property, the top layer maps the inputs into a low-dimensional feature space . Given the teacher graph, SNTG makes the learned features more discriminative by enforcing them to be similar for neighbors and dissimilar for those non-neighbors. The model structure is depicted in Fig. 1. We then propose a doubly stochastic sampling algorithm to reduce the computational cost with large mini-batch sizes. Our method can be applied with very little engineering effort to existing deep SSL works including both generative and discriminative approaches because SNTG does not introduce any extra network parameters. We demonstrate significant performance improvements over state-of-the-art results while the extra time cost is negligible.

Related work

Using unlabeled data to improve generalization has a long and rich history and the literature in SSL is vast . So in this section we focus on reviewing the closely related papers, especially the recent advances in SSL with deep learning.

Self-training methods iteratively use the current classifier to label those unlabeled ones with high confidence . Co-training uses a pair of classifiers with disjoint views of data to iteratively learn and generate training labels. Transductive SVMs implement the cluster assumption by keeping unlabeled data far away from the decision boundaries. Entropy minimization , a strong regularization term commonly used, minimizes the conditional entropy H(p(y∣x))H\left(p\left(y|x\right)\right) to ensure that one instance is assigned to one class with a high probability to avoid class overlap.

Graph-based Methods. Graph-based SSL methods define the similarity of data points by a graph and make predictions smooth with respect to the graph structure. Many of them often optimize a supervised loss over labeled data with a graph Laplacian regularizer . Label propagation pushes label information from a labeled instance to its neighbors using a predefined distance metric. We emphasize that our work differs from these traditional methods in the construction and utilization of the graph. Previous work usually constructs the graph in advance using prior knowledge or manual labeling and the graph remains fixed in the following training process . This can lead to several disadvantages as detailed in Sec. 4.2 and 5.3. Although some works establish the graph dynamically during the classification, their performance is far from recent state-of-the-art deep learning based methods.

Generative Approaches. Besides aforementioned discriminative approaches, another line is generative models, which pay efforts to learn the input distribution p(x)p(x) that is believed to share some information with the conditional distribution p(y∣x)p(y|x) . Traditional models such as Gaussian mixtures try to maximize the joint log-likelihood of both labeled and unlabeled data using EM. For modern deep generative models, variational auto-encoder (VAE) makes it scalable by employing variational methods combined with deep learning while generative adversarial networks (GAN) generate samples by optimizing an adversarial game between the discriminator and the generator . The samples generated by GAN can be viewed as another kind of “data augmentation” to “tell” the decision boundary where to lie. For example, “fake” samples can be generated in low density regions where the training data is rare based on the low density separation assumption. Alternatively, more “pseudo” samples could be generated in high density regions to keep away from the decision boundary thus improve the robustness of the classifier . Our work is complementary to these efforts and can be easily combined with them. We observe improvements over feature matching GAN with SNTG (see Section 5.6).

Background

We consider the semi-supervised classification task, where the training set D\mathcal{D} consists of NN examples, out of which LL have labels and the others are unlabeled. Let L={(xi,yi)}i=1L\mathcal{L}=\{(x_{i},y_{i})\}_{i=1}^{L} be the labeled set and U={xi}i=L+1N\mathcal{U}=\{x_{i}\}_{i=L+1}^{N} be the unlabeled set where the observation xi∈Xx_{i}\in\mathcal{X} and the corresponding label yi∈Y={1,2,...,K}y_{i}\in\mathcal{Y}=\{1,2,...,K\}. We aim to learn a function f:X→Kf:\mathcal{X}\to^{K} parameterized by θ∈Θ\theta\in\Theta by solving a generic optimization problem:

As mentioned earlier, the models in perturbation-based methods assume a dual role, i.e., a teacher and a student . The training targets for the student are generated by the teacher. Recent progresses focus on improving the quality of targets by using self-ensembling and exploring different perturbations , as summarized in . Formally, self-ensembling methods fit in Eq. (1) by defining RR as a consistency loss:

Mean teacher (MT) . Instead of averaging predictions every epoch, MT updates the targets more frequently to form a better teacher, i.e., it averages parameters θ\theta every iteration:

Virtual adversarial training (VAT) . Instead of l2l_{2} distance, VAT defines RR as the KL divergence between the model prediction and that of the input under adversarial perturbations ξadv′\xi^{\prime}_{adv}:

Our approach

One common shortcoming of the perturbation-based methods is that they regularize the output to be smooth near a data point locally, while ignoring the cluster structure. We address it by proposing a new SSL method, SNTG, that enforces neighbors to be smooth, which is a stronger regularization than only imposing smoothness at a single unlabeled point. We show that SNTG contributes to form a better teacher model, which is the focus of recent advances on perturbation-based methods. In the following, we formalize our approach by answering two key questions: (1) how to define the graph and neighbors? and (2) how to induce the smoothness of neighboring points using the graph?

Most existing graph-based SSL methods depend on a distance metric in the input space X\mathcal{X}, which is typically low-level (e.g., pixel values of images). For natural images, pixel distance cannot reflect semantic similarity well. Instead, we use the distance in the label space Y\mathcal{Y}, and treat the data points from the same class as neighbors. However, an issue is that the true labels of unlabeled data are unknown. We address it by learning a teacher graph using the targets generated by the teacher model. Self-ensembling is a good choice for constructing the graph because the ensemble predictions are expected to be more accurate than the outputs of current classifier. Inspired by that, a teacher graph can guide the student model to move in correct directions. A comparison to other graphs could be found in Sec. 5.3.

2 Guiding the low-dimensional feature mapping

Given a N×NN\times N similarity matrix WW of the sparse graph, we define the SNTG loss as

where m>0m>0 is a pre-defined margin and ∥⋅∥\|\cdot\| is Euclidean distance. The margin loss is to constrain neighboring points to have consistent features. Consequently, the neighbors are encouraged to have consistent predictions while the non-neighbors (i.e., the points of different classes) are pushed apart from each other with a minimum distance mm. Visualizations can be found in Section 5.4.

We discuss the difference between SNTG and two early works LPDGL and EmbedNN . For LPDGL, the definition and the usage of local smoothness are both different from ours. LPDGL defines deformed Laplacian to smooth the predictions of kk neighbors in a local region while our work enforces the features to be smooth by the contrastive loss in Eq. (8) w.r.t. the 0-1 teacher graph. For EmbedNN, despite they also measure the embedding loss, there are several key differences. First, inspired by Π\Pi model, SNTG aims to induce more smoothness using neighbors under perturbations, while EmbedNN is motivated by using the embedding as an auxiliary task to help supervised tasks and does not consider the robustness to perturbations. Second, EmbedNN uses a fixed graph WW defined by kk-nearest-neighbor (kk-NN) based on the distance in X\mathcal{X}. Our method takes a different approach using the teacher-generated targets in Y\mathcal{Y}. As mentioned in Section 4.1, the pixel-level distance in X\mathcal{X} may not reflect the semantic similarity as well as that in Y\mathcal{Y} for natural images. Third, once the graph is built in EmbedNN, the fixed graph cannot leverage the knowledge distilled by the classifier thus cannot be improved any more, while SNTG jointly learns the classifier and the teacher graph as stated above. Furthermore, on the time cost and scalability, SNTG is faster than EmbedNN and can handle large-scale datasets. kk-NN in EmbedNN is slow for large kk and even more time-consuming for large-scale datasets. We compute WW in the much lower dimensional Y\mathcal{Y} and use the sub-sampling technique that is to be introduced next. Experimental comparisons are in Section 5.3.

3 Doubly stochastic sampling approximation

Our overall objective is the sum of two components. The first one is the standard cross-entropy loss on the labeled data, and the second is the regularization term, which encourages the smoothness for each single point (i.e., RCR_{C}) as well as for the neighboring points (i.e., RSR_{S}). Alg. 1 presents the pseudo-code. Following , we use a ramp-up w(t)w(t) for both the learning rate and the regularization term in the beginning.

As our model uses deep networks, we train it using Stochastic Gradient Descent (SGD) with mini-batches. We follow the common practice and construct the sub-graph in a random mini-batch to estimate RSR_{S} in Eq. (7). For a mini-batch BB of size nn, we need to compute WijW_{ij} for all the data pairs (xi,xj)∈B(x_{i},x_{j})\in B, which is of size n2n^{2} in total. Although this step is fast, the computation of ∥h(xi)−h(xj)∥\|h(x_{i})-h(x_{j})\| related to WijW_{ij} is O(p)O(p) and then the overall computational cost is O(n2p)O(n^{2}p), which is slow for large nn. To reduce the computational cost, we instead use doubly stochastic sampled data pairs to construct WijW_{ij} and only use them to compute Eq. (8), which is still an unbiased estimation of RSR_{S}. Specifically, in each iteration, we sample a mini-batch BB and then sub-sample s≤n2s\leq n^{2} data pairs SS from BB. Empirically, SNTG can be incorporated into other SSL methods with not much extra time cost. See Appendix A for details.

Experiments

This section presents both quantitative and qualitative results to demonstrate the effectiveness of SNTG. The purpose of experiments is to show the improvements that come from SNTG, using cutting-edge approaches as evidence. Source code is at https://github.com/xinmei9322/SNTG.

2 Benchmark datasets

We then provide results on the widely adopted benchmarks, MNIST, SVHN, CIFAR-10 and CIFAR-100. Following common practice , we randomly sample 100, 1000 4000 and 10000 labels for MNIST, SVHN, CIFAR-10 and CIFAR-100, respectively. We further explore fewer labels for the non-augmented MNIST as well as SVHN and CIFAR-10 with standard augmentation. The results are averaged over 10 runs with different seeds for data splits. Main results are presented in Tables 1, 2, 3 and 4. The accuracy of baselines are all taken from existing literature. In general, we can see that our method surpasses previous state-of-the-arts by a large margin.

All models are trained with the same network architecture and hyper-parameters to our baselines, i.e., perturbation-based methods described in Sec. 3.1. The SNTG loss only needs three extra hyper-parameters: the regularization parameter λ2\lambda_{2}, the margin mm and the number of sub-sampled pairs ss. We fix mm and ss, and only tune λ2\lambda_{2}. More details on experimental setup can be found in Appendix A. For fair comparison, we also report our best implementation under the settings not covered in (marked ∗*).

Note that VAT is a much stronger baseline than Π\Pi model and TempEns since it explores adversarial perturbation with extra efforts and more time. VAT’s best results are achieved with an additional entropy minimization (Ent) regularization . We evaluate our method under the best setting VAT+Ent and observe a further improvement with SNTG, e.g., from 13.15% to 12.49% and from 10.55% to 9.89% on CIFAR-10 without or with augmentation, respectively. In fact, we observed that Ent could also improve the performance of other self-ensembling methods if it was added along with SNTG. But to keep the results clear and focus on the efficacy of SNTG, we did not illustrate the results here.

As shown in Tables 2 and 3, when SNTG is applied to the fully supervised setting (i.e., all labels are observed), our method further reduces the error rates compared to self-ensembling methods, e.g., from 5.56%5.56\% to 5.19%5.19\% on CIFAR-10 for Π\Pi model. It suggests that supervised learning also benefits from the additional smoothness and the learned invariant feature space in SNTG.

Fewer labels. Notably, as shown in Tables 4, 2 and 3, when labels are very scarce, e.g., MNIST with 20 labels (only 2 labeled samples per class), SVHN with 250 labels and CIFAR-10 with 1000 labels, the benefits provided by SNTG are even more significant. The SNTG regularizer empirically reduces the overfitting on the small set of labeled data and thus yields better generalization.

Ablation study. Our reported results are based on adding SNTG loss RSR_{S} to baselines, and the overall objective has already included the consistency loss RCR_{C} (See Alg. 1, line 9). To quantify the effectiveness of our method, Table 5 presents the evaluation of Π\Pi+SNTG compared to its ablated versions. The error rate of Π\Pi model, which only uses RCR_{C}, is 16.55%16.55\%. However, using RSR_{S} alone yields a lower error rate of 15.36%15.36\%. Thus, RSR_{S} considering the neighbors proves to be a strong regularization, comparable or even favorable to RCR_{C}, and they are also complementary.

Convergence. A potential concern of our method is the convergence, since the information in a teacher graph is likely to be inaccurate at the beginning of training. However, we did not observe any divergent cases in all experiments. Empirically, the teacher model is usually a little better than the student in training. Furthermore, the ramp-up w(t)w(t) is used to balance the trade-off between the supervised loss and regularization, which is important for the convergence as described in previous works . Using the ramp-up weighting mechanism, the supervised loss dominates the learning in earlier training. As the training continues, the student model has more confidence in the information given by the teacher model, i.e., the target predictions and the graph, which gradually contributes more to the learning process. Fig. 3 shows that our model converges well.

3 Comparison to EmbedNN and other graphs

As our graph is learned based on the predictions in Y\mathcal{Y} given by the teacher model, we further compare to other graphs. We test them on CIFAR-10 using 4000 labels without augmentation and share all the same hyper-parameter settings with Π\Pi model except the definition of WW. The first baseline is a fixed graph defined by kk-NN in X\mathcal{X}—Following EmbedNN , WW is predefined so that 10 nearest neighbors of xix_{i} have Wij=1W_{ij}=1, and Wij=0W_{ij}=0 otherwise. The second one is another fixed graph in Y\mathcal{Y}—Since only a small portion of labels are observed on training data in SSL, we construct the graph based on the predictions of a pre-trained Π\Pi model on training data. Fig. 3 shows that our model outperforms other graphs. The test error rate of the baseline Π\Pi model is 16.55%16.55\%. Using kk-NN in X\mathcal{X} gives a marginal improvement to 16.13%16.13\%. Using the predictions in pre-trained Π\Pi model to construct a 0-1 fixed graph, the error rate is 15.71%15.71\%. Using our method, learning a teacher graph from scratch, Π\Pi+SNTG achieves superior result with 13.62%13.62\% error rate.

Note that Π\Pi model is a strong baseline surpassing most previous methods. For natural images like CIFAR-10, the pixel-level distance provides limited information for the similarity thus kk-NN graph in X\mathcal{X} does not improve the strong baseline. The reason of the performance gap to the second one lies in that using a fixed graph in Y\mathcal{Y} is more like “pre-training” while using teacher graph is like “joint-training”. The teacher graph becomes better using the information extracted by the teacher and then benefits it in turn. However, the fixed graphs cannot receive feedbacks from the model in the training and all the information is from the pre-training or prior knowledge. Empirical results support our analysis.

4 Visualization of embeddings

5 Robustness to noisy labels

We finally show that SNTG can not only benefit from unlabeled data, but also learn from noisy supervision. Following , we did extra experiments on supervised SVHN to show the tolerance to incorrect labels. Certain percentages of true labels on the training set are replaced by random labels. Fig. 4 shows that TempEns+SNTG retains over 93% accuracy even when 90% of the labels are noisy while TempEns alone only obtains 73% accuracy . With standard supervised training, the model suffers a lot and overfits to the incorrect information in labels. Thus, our SNTG regularization improves the robustness and generalization performance of the model. Previous work also shows that self-generated targets yield robustness to label noise.

6 Feature matching GAN benefits from SNTG

Recently, the feature matching (FM) GAN in Improved GAN has performed well for SSL but usually generates images with strange patterns. Some works have been done to analyze the reasons . An interesting finding is that our method can also alleviate the problem. Fig. 6 presents the comparison between the samples generated in FM GAN and FM GAN+SNTG. Apart from improving the generated sample quality of FM GAN, SNTG also reduces the error rate. FM GAN achieves 18.63%18.63\% on CIFAR-10 with 4000 labels. We regularize the features of unlabeled data using SNTG and observe an improvement to 14.93%14.93\%, which is comparable to the state-of-the-art 14.41%14.41\% in deep generative models .

In FM GAN, the objective for the generator is defined as

which is similar to the neighboring case when Wij=1W_{ij}=1 in Eq. (8). In our opinion, SNTG helps shape the feature space better so that the generator could capture the data distribution by matching only the mean of features.

Conclusions and future work

We present a simple but effective SNTG, which regularizes the neighboring points on a learned teacher graph. Empirically, it outperforms all baselines and achieves new state-of-the-art results on several datasets. As a byproduct, we also learn an invariant mapping on a low-dimensional manifold. SNTG offers additional benefits such as handling extreme cases with fewer labels and noisy labels. In future work, it is promising to do more theoretical analysis of our method and to explore its combination with generative models as well as applications to large-scale datasets, e.g., ImageNet with more classes.

Acknowledgements

The work is supported by the National NSF of China (Nos. 61620106010, 61621136008, 61332007), Beijing Natural Science Foundation (No. L172037), Tsinghua Tiangong Institute for Intelligent Computing, the NVIDIA NVAIL Program and a research fund from Siemens.

References

Appendix A Experimental setup

MNIST. It contains 60,000 gray-scale training images and 10,000 test images from handwritten digits to 99. The input images are normalized to zero mean and unit variance.

SVHN. Each example in SVHN is a 32×3232\times 32 color house-number images and we only use the official 73,257 training images and 26,032 test images following previous work. The augmentation of SVHN is limited to random translation between $$ pixels.

CIFAR-10. The CIFAR-10 dataset consists of 32×3232\times 32 natural RGB images from 10 classes such as airplanes, cats, cars and horses. We have 50,000 training examples and 10,000 test examples. The input images are normalized using ZCA following previous work . We use the standard way of augmenting the CIFAR-10 dataset including horizontal flips and random translations.

CIFAR-100. The CIFAR-100 dataset consists of 32×3232\times 32 natural RGB images from 100 classes. We have 50,000 training examples and 10,000 test examples. The preprocession of inputs images are the same to CIFAR-10.

Implementation. We implemented our code mainly in Python with Theano and Lasagne . For comparison with VAT and Mean Teacher experiments, we use TensorFlow to match their settings. The code for reproducing the results is available at https://github.com/xinmei9322/SNTG.

Training details. In Π\Pi model and TempEns based experiments, the network architectures (shown in Table 6) and the hyper-parameters are the same as our baselines . We apply mean-only batch normalization with momentum 0.9990.999 to all layers and use leaky ReLU with α=0.1\alpha=0.1. The network is trained for 300300 epochs using Adam Optimizer with mini-batches of size n=100n=100 and maximum learning rate 0.0030.003 (exceptions are that TempEns for SVHN uses 0.0010.001 and MNIST uses 0.00010.0001). We use the default Adam momentum parameters β1=0.9\beta_{1}=0.9 and β2=0.999\beta_{2}=0.999. Following , we also ramp up the learning rate and the regularization term during the first 80 epochs with weight w(t)=exp⁡[−5(1−t80)2]w(t)=\exp\left[-5(1-\frac{t}{80})^{2}\right] and ramp down the learning rate during the last 50 epochs. The ramp-down function is exp⁡[−12.5(1−300−t50)2]\exp\left[-12.5(1-\frac{300-t}{50})^{2}\right]. The regularization coefficient of consistency loss RCR_{C} is λ1=100\lambda_{1}=100 for Π\Pi model and λ1=30\lambda_{1}=30 for TempEns (exception is that SVHN with L=250L=250 uses λ1=50\lambda_{1}=50).

For comparison with Mean Teacher and VAT, we keep the same architecture and hyper-parameters settings with the corresponding baselines . Their network architectures are the same as shown in Table 6 but differ in several hyper-parameters such as weight normalization, training epochs and mini-batch sizes, which are detailed in their papers. We just add the SNTG loss along with their regularization RCR_{C} and keep other settings unchanged as in their public code.

Training time. SNTG does not increase the number of neural network parameters and the runtime is almost the same to the baselines, with only extra 1-2 seconds per epoch (the baselines usually need 100-200 seconds per epoch on one GPU).

Synthetic benchmarks. The synthetic dataset experiments adopt the default settings for Π\Pi model except for 0.0010.001 maximum learning rate and 500500 training epochs. We use weight normalization and add Gaussian noise to each layer.

Appendix B Rethinking ΠΠ\Pi model objective

In Π\Pi model , the consistency loss is defined in Eq. (2) where the teacher model shares the same parameter with the student model θ′=θ\theta^{\prime}=\theta. Suppose f(x)∈Kf(x)\in^{K}, the consistency loss of Π\Pi model is

where [⋅]k[\cdot]_{k} is kk-th component of the vector.

Then minimizing RCR_{C} is equivalent to minimizing the sum of variance of the prediction each dimension. Similar idea of variance penalty was exploited in Pseudo-Ensemble . If a data point is near the decision boundary, it is likely to has a large variance since its prediction might alternate to another class when some noise is added. Minimizing the variance explicitly penalizes such alternation behavior of training data.

Appendix C Comparison to classical SSL methods

As mentioned in Section 2, our method is different from classical graph-based SSL methods in many important aspects such as the construction of the graph and how to use it.

Table 7 is a comparion with several classical methods: (1) Label propagation (LP) ; (2) A variant of LP on kkNN structure(LP+kkNN) ; (3) Local and Global Consistency (LGC) ; (5) Transductive SVM (TSVM) ; (6) LapRLS ; (7) Dynamic Label propagation (DLP) . The results of (1)-(7) are cited from . We also compare with the best reported results in previously mentioned works: (8) EmbedNN ; (9) the Manifold Tangent Classifier (MTC) ; (10) Pseudo-Ensemble .

While the classical graph-based methods (e.g., LP, DLP and LapRLS) were the leading paradigms, with the resurgence of deep learning, recent impressive results are mostly from deep learning based SSL methods, while classical methods fall behind on performance and scalability. Furthermore, they have no reported results on challenging natural image datasets, e.g., SVHN, CIFAR-10. Only one overlap is MNIST, see Table 7 for comparison. We show that our method SNTG surpasses these classical methods by a large margin.

Appendix D Significance test of the improvements.

Table 8 shows the independent two sample T-test on the error rates of baselines and our method. All the P-values are less than significance level α=0.01\alpha=0.01. It indicates that the improvements of SNTG are significant.