FlexMatch: Boosting Semi-Supervised Learning with Curriculum Pseudo Labeling
Bowen Zhang, Yidong Wang, Wenxin Hou, Hao Wu, Jindong Wang, Manabu Okumura, Takahiro Shinozaki
Introduction
Semi-supervised learning (SSL) has attracted increasing attention in recent years due to its superiority in leveraging a large amount of unlabeled data. This is particularly advantageous when the labeled data are limited in quantity or laborious to obtain. Consistency regularization and pseudo labeling are two powerful techniques for utilizing unlabeled data and have been widely used in modern SSL algorithms . The recently proposed FixMatch achieves competitive results by combining these techniques with weak and strong data augmentations and using cross-entropy loss as the consistency regularization criterion.
However, a drawback of FixMatch and other popular SSL algorithms such as Pseudo-Labeling and Unsupervised Data Augmentation (UDA) is that they rely on a fixed threshold to compute the unsupervised loss, using only unlabeled data whose prediction confidence is above the threshold. While this strategy can make sure that only high-quality unlabeled data contribute to the model training, it ignores a considerable amount of other unlabeled data, especially at the early stage of the training process, where only a few unlabeled data have their prediction confidence above the threshold. Moreover, modern SSL algorithms handle all classes equally without considering their different learning difficulties.
To address these issues, we propose Curriculum Pseudo Labeling (CPL), a curriculum learning strategy to take into account the learning status of each class for semi-supervised learning. CPL substitutes the pre-defined thresholds with flexible thresholds that are dynamically adjusted for each class according to the current learning status. Notably, this process does not introduce any additional parameter (hyperparameter or trainable parameter) or extra computation (forward or back propagation). We apply this curriculum learning strategy directly to FixMatch and call the improved algorithm FlexMatch.
While the training speed remains as efficient as that of FixMatch, FlexMatch converges significantly faster and achieves state-of-the-art performances on most SSL image classification benchmarks. The benefit of introducing CPL is particularly remarkable when the labels are scarce or when the task is challenging. For instance, on the STL-10 dataset, FlexMatch achieves relative performance improvement over FixMatch by 18.96%, 16.11%, and 7.68% when the label amount is 400, 2500, and 10000 respectively. Moreover, CPL further shows its superiority by boosting the convergence speed – with CPL, FlexMatch takes less than 1/5 training time of FixMatch to reach its final accuracy. Adapting CPL to other modern SSL algorithms also leads to improvements in accuracy and convergence speed.
To sum up, this paper makes the following three contributions:
We propose Curriculum Pseudo Labeling (CPL), a curriculum learning approach of dynamically leveraging unlabeled data for SSL. It is almost cost-free and can be easily integrated to other SSL methods.
CPL significantly boosts the accuracy and convergence performance of several popular SSL algorithms on common benchmarks. Specifically, FlexMatch, the integration of FixMatch and CPL, achieves state-of-the-art results.
We open-source TorchSSL, a unified PyTorch-based semi-supervised learning codebase for the fair study of SSL algorithms. TorchSSL includes implementations of popular SSL algorithms and their corresponding training strategies, and is easy to use or customize.
Background
where is the batch size of labeled data, is the ratio of unlabeled data to labeled data, is a stochastic data augmentation function (thus the two terms in Eq.(1) are different), denotes a piece of unlabeled data, and represents the output probability of the model. With the introduction of pseudo labeling techniques , the consistency regularization is converted to an entropy minimization process , which is more suitable for the classification task. The improved consistency loss with pseudo labeling can be represented as:
where is cross-entropy, is the pre-defined threshold and is the pseudo label that can either be a ‘hard’ one-hot label or a sharpened ‘soft’ one . The intention of using a threshold is to mask out noisy unlabeled data that have low prediction confidence.
FixMatch utilizes such consistency regularization with strong augmentation to achieve competitive performance. For unlabeled data, FixMatch first uses weak augmentation to generate artificial labels. These labels are then used as the target of strongly-augmented data. The unsupervised loss term in FixMatch thereby has the form:
where is a strong augmentation function instead of weak augmentation .
Of the aforementioned works, the pre-defined threshold () is constant. We believe this can be improved because the data of some classes may be inherently more difficult to learn than others. Curriculum learning is a learning strategy where learning samples are gradually introduced according to the model’s learning process. In such a way, the model is always optimally challenged. This technique is widely employed in deep learning research .
FlexMatch
While current SSL algorithms render pseudo labels of only high-confidence unlabeled data cut off by a pre-defined threshold, CPL renders the pseudo labels to different classes and at different time steps. Such a process is realized by adjusting the thresholds according to the model’s learning status of each class.
However, it is non-trivial to dynamically determine the thresholds according to the learning status. The most ideal approach would be calculating evaluation accuracies for each class and use them to scale the threshold, as:
where is the flexible threshold for class at time step and is the corresponding evaluation accuracy. In this way, lower accuracy that indicates a less satisfactory learning status of the class will lead to a lower threshold that encourages more samples of this class to be learned. Since we cannot use the evaluation set in the model learning process, one may have to separate an extra validation set from the training set for such accuracy evaluations. However, this practice show two fatal problems: First, such a labeled validation set separated from the training set is expensive under SSL scenario as the labeled data are already scarce. Second, to dynamically adjust the thresholds in the training process, accuracy evaluations must be done continually at each time step , which will considerably slow down the training speed.
In this work, we propose Curriculum Pseudo Labeling (CPL) for semi-supervised learning. Our CPL uses an alternative way to estimate the learning status, which does not introduce additional inference processes, nor needs an extra validation set. As believed in , a high threshold that filters out noisy pseudo labels and leaves only high-quality ones can considerably reduce the confirmation bias . Therefore, our key assumption is that when the threshold is high, the learning effect of a class can be reflected by the number of samples whose predictions fall into this class and above the threshold. Namely, the class with fewer samples having their prediction confidence reach the threshold is considered to have a greater learning difficulty or a worse learning status, formulated as:
where reflects the learning effect of class at time step . is the model’s prediction for unlabeled data at time step , and is the total number of unlabeled data. When the unlabeled dataset is balanced (i.e., the number of unlabeled data belonging to different classes are equal or close), larger indicates a better estimated learning effect. By applying the following normalization to to make its range between to , it can then be used to scale the fixed threshold :
One characteristic of such a normalization approach is that the best-learned class has its equal to , causing its flexible threshold equal to . This is desirable. For classes that are hard to learn, the thresholds are lowered down, encouraging more training samples in these classes to be learned. This also improves the data utilization ratio. As learning proceeds, the threshold of a well-learned class is raised higher to selectively pick up higher-quality samples. Eventually, when all classes have reached reliable accuracies, the thresholds will all approach . Note that the thresholds do not always grow, it may also decrease if the unlabeled data is classified into a different class in later iterations. This new threshold is used for calculating the unsupervised loss in FlexMatch, which can be formulated as:
where . The flexible thresholds are updated at each iteration. Finally, we can formulate the loss in FlexMatch as the weighted combination (by ) of supervised and unsupervised loss:
where is the supervised loss on labeled data:
Note that the cost of introducing CPL is almost free. Practically, every time the prediction confidence of an unlabeled data is above the fixed threshold , the data, and its predicted class are marked and will be used for calculating at the next time step. Such marking actions are bonus actions each time the consistency loss is computed. Therefore, FlexMatch does not introduce additional forward propagation processes for evaluating the model’s learning status, nor new parameters.
2 Threshold warm-up
We noticed in our experiments that at the early stage of the training, the model may blindly predict most unlabeled samples into a certain class depending on the parameter initialization(i.e., more likely to have confirmation bias). Hence, the estimated learning status may not be reliable at this stage. Therefore, we introduce a warm-up process by rewriting the denominator in Eq. (6) as:
where the term can be regarded as the number of unlabeled data that have not been used. This ensures that at the beginning of the training, all estimated learning effects gradually rise from until the number of unused unlabeled data is no longer predominant. The duration of such a period depends on the unlabeled data amount (ref. in Eq. (11)) and the learning difficulty (ref. the growing speed of in Eq. (11)) of the dataset. In practice, such a warm-up process is very easy to implement as we can add an extra class to denote the unused unlabeled data. Thus calculating the denominator of Eq. (11) is simply converted to finding the maximum among classes.
3 Non-linear mapping function
The flexible threshold in Eq. (7) is determined by the normalized estimated learning effects via a linear mapping. However, it may not be the most suitable mapping in the real training process, where the increase or decrease of may make big jumps in the early phase where the predictions of the model are still unstable; and only make small fluctuations after the class is well-learned in the mid and late training stage. Therefore, it is preferable if the flexible thresholds can be more sensitive when is large and vice versa.
We propose a non-linear mapping function to enable the thresholds to have a non-linear increasing curve when ranges uniformly from to , as formulated below:
where is a non-linear mapping function. It is clear that Eq. (7) can be seen as a special case by setting to the identity function. The mapping function should be monotonically increasing and have a maximum no larger than (otherwise the flexible threshold can be larger than and filter out all samples). To avoid introducing additional hyperparameters (e.g. lower limits of the flexible thresholds), we consider the mapping function to have a range from to so that the flexible thresholds range from to .
A monotone increasing convex function lets the thresholds grow slowly when is small, and become more sensitive as gets larger. Hence, we intuitively choose a convex function with the above-mentioned properties for our experiments. We also conduct an ablation study to compare among mapping functions with different convexity and concavity in Sec. 4.4. The full algorithm of FlexMatch is shown in Algorithm 1.
Experiments
We evaluate FlexMatch and other CPL-enabled algorithms on common SSL datasets: CIFAR-10/100 , SVHN , STL-10 and ImageNet , and extensively investigate the performance under various labeled data amounts. We mainly compare our method with Pseudo-Labeling , UDA and FixMatch , since they all involve a pre-defined threshold. The results of other popular SSL algorithms are in the appendix B. We also add a fully-supervised experiment for each dataset to better understand the results of SSL algorithms. Note that previously suggested fully-supervised comparisons use only the labeled set for training, whose purpose is to manifest the improvement brought by the introduction of unlabeled data. With the development of modern SSL algorithms, however, semi-supervised approaches are achieving competitive performance with supervised ones, or even better performance due to the strength of consistency regularization. Therefore, our fully-supervised comparisons are conducted with all data labeled, and apply weak data augmentations following Eq. (10). We re-implement all baselines using our PyTorch codebase: TorchSSL, which is introduced in the appendix B.
For a fair comparison, we use the same hyperparameters following FixMatch . Concretely, the optimizer for all experiments is standard stochastic gradient descent (SGD) with a momentum of 0.9 . For all datasets, we use an initial learning rate of with a cosine learning rate decay schedule as , where is the initial learning rate, is the current training step and is the total training step that is set to . We also perform an exponential moving average with the momentum of . The batch size of labeled data is except for ImageNet. is set to be for Pseudo-Label and for UDA, FixMatch, and FlexMatch. is set to for UDA and for Pseudo Label, FixMatch, and FlexMatch. These setups follow the original papers. The strong augmentation function used in our experiments is RandAugment . We use ResNet-50 for the ImageNet experiment and Wide ResNet (WRN) and its variant for other datasets. Detailed hyperparameters are listed in the appendix A.
We adopt two evaluation metrics: (1) the median error rate of the last 20 checkpoints following , and (2) the best error rate in all checkpoints. We argue that the median approach is not suitable when the convergence speeds of the algorithms show significant differences – the large number of redundant iterations may result in over-fitting for the fast-converge algorithms. Therefore, we report the best error rates for all algorithms, while the results of the median approach are also provided in the appendix A, showing that our FlexMatch still achieves the best performance. We run each task three times using distinct random seeds to obtain the error bars.
The classification error rates on CIFAR-10/100, STL-10 and SVHN datasets are in Table 1, and the results on ImageNet are in Sec. 4.2. Note that the SVHN dataset used in our experiment also includes the extra set that contains 531,131 additional samples. Results demonstrate that FlexMatch achieves the state-of-the-art performance on most of the benchmark datasets except for SVHN where Flex-UDA (i.e., UDA with CPL) and UDA have the lowest error rate on the 40-label split and the 1000-label split, respectively. We also provide the detailed precision, recall, F1, and AUC results in the appendix A. Our CPL (FlexMatch) has the following advantages:
Our FlexMatch significantly outperforms other methods when the amount of labels is extremely small. For instance, on the CIFAR-100 dataset with 400 labels (i.e., only 4 label samples per class), FlexMatch achieves an average error rate of 39.94%, which significantly outperforms FixMatch (46.42%).
CPL improves the performance of existing SSL algorithms.
Other than FixMatch, CPL can also improve the performance of other existing SSL algorithms such as Pseudo-Labeling and UDA. For instance, the error rate is reduced from 37.4% to 29.53% for UDA on the STL-10 40-label split after introducing CPL (refer to as Flex-UDA in Table 1). These results further prove the effectiveness of CPL in better leveraging unlabeled data. Figure 2 shows the average running time of a single iteration with or without adding our CPL, it is clear that while improving the performance of existing SSL algorithms, our CPL does not introduce additional computational burden.
CPL achieves better performance on complicated tasks.
The STL-10 dataset contains unlabeled data from a similar but broader distribution of images than its labeled set. The existence of new types of objects in the unlabeled dataset makes STL-10 a more challenging and realistic task. FlexMatch achieves greater performance improvement under such a challenging situation. The error rate on STL-10 with only 40 labels is 29.15%, which is relatively 18.96% better than FixMatch (35.97%). Similar strong improvements are also observed on CIFAR-100 dataset, which has as many as one hundred classes.
We also analyze the reason why FlexMatch performs less favorably on SVHN. This is probably because SVHN is a relatively simple (i.e., to classify digits) yet unbalanced dataset. The class-wise imbalance leads to the classes with fewer samples never have their estimated learning effects close to 1 according to Eq. (6), even when they are already well-learned. Such low thresholds allow noisy pseudo-labeled samples to be trusted and learned throughout the training process, which is also reflected by the loss descent curve where the low-threshold classes have major fluctuations. FixMatch, on the other hand, fixes its threshold at 0.95 to filter out noisy samples. Such a fixed high threshold is not preferable with respect to both accuracies of hard-to-learn classes and overall convergence speed as explained earlier, but since SVHN is an easy task, the model can easily learn the task and make high-confidence predictions, setting a high-fixed threshold thus becomes less problematic and has its advantages overweighed.
2 Results on ImageNet
We also verify the effectiveness of CPL on ImageNet-1K which is a much more realistic and complicated dataset. We randomly choose the same 100K labeled data (i.e., 100 labels per class), which is less than of the total labels. The hyper-parameters used for ImageNet can be found in the appendix A, where the two algorithms share the same hyper-parameters. We show the error rate comparison after running iterations in Table 2. This result indicates that when the task is complicated, despite the class imbalance issue (the number of images within each class ranges from 732 to 1300), CPL can still bring improvements. Note that this result does not represent the best performance of each algorithm as the model cannot fully converge after iterations, and due to the computational resource limitation, we did not further tune the hyper-parameters to obtain the best results on ImageNet.
3 Convergence speed acceleration
Another strong advantage of FlexMatch is its superior convergence speed. Figure 3(a) and 3(b) shows the comparison between FlexMatch and FixMatch with respect to the loss and top-1-accuracy on CIFAR-100 400-label split. The loss of FlexMatch decreases much faster and smoother than FixMatch, demonstrating its superior convergence speed. The major fluctuations of the loss in FixMatch may due to the pre-defined threshold that lets pass most unlabeled data belonging to certain classes, whereas with CPL a larger batch of unlabeled data containing samples from various classes enables the gradient to more directly head toward the global optimum. As a result, with only 50K iterations, FlexMatch has already surpassed the final results of FixMatch. After 800K iterations, however, we observe a further decrease in loss and accuracy. This is likely due to over-fitting, which also occurs in FixMatch after 900K iterations. Thus, we believe it is not fair to use the median results of the last few checkpoints for evaluating algorithms with different convergence speeds.
We further compare the class-wise accuracy of FixMatch and FlexMatch on CIFAR-10 in their early training stages. As shown in Figure 3(c) and 3(d), at iteration 200K, FixMatch only hits an overall accuracy of 56.35% as half of the classes are still learned unsatisfactorily, whereas FlexMatch has already achieved an overall accuracy of 94.29% which is even higher than the final accuracy reached by FixMatch after 1M iterations. It is manifest that the introduction of CPL successfully encourages the model to proactively learn those difficult classes thereby improving the overall learning effect.
4 Ablation study
We conduct experiments to evaluate three components of FlexMatch: the upper limit of thresholds , mapping functions , and threshold warm-up.
We investigate 5 different values and 3 different mapping functions on CIFAR-10 dataset with 40 labels. As shown in Figure 4(a), the optimal choice of is around , either increasing or decreasing this value results in a performance decay. Note that in FlexMatch, tuning does not only affect the upper limit of the threshold but also the estimated learning effects because they are determined by the number of samples that fall above .
Mapping function.
We explore three different mapping functions in Figure 4(b): (1) concave: , (2) linear: , and (3) convex: . We see that the convex function shows the best performance and the concave function shows the worst. Although tweaking the degree of convexity may probably lead to further improvement, we do not make further investigation in this paper. It is noteworthy that all these functions have their outputs grow from to when the inputs go from to . One may also design a function with a different range, for instance, from to . In this case, it is equivalent to setting a lower limit to the flexible threshold so that even at the beginning of the training, only samples with prediction confidence higher than this limit will contribute to the unsupervised loss. We do not include such a lower limit in FlexMatch since it will introduce a new hyperparameter. However, we did find that setting a lower limit at can slightly improve the performance. A possible reason is that the lower threshold prevents noisy training caused by incorrect pseudo labels at the early stage .
Threshold warm-up.
We analyze the performance of threshold warm-up on both CIFAR-10 (40 labels) and CIFAR-100 (400 labels) datasets. As shown in Figure 4(c), threshold warm-up can bring about 0.2% absolute improvement on CIFAR-10 and about 1% on CIFAR-100. At the beginning of the training without the threshold warm-up, the flexible thresholds may go through heavy fluctuations because the denominator in Eq.(6) is small. In the meantime, there will always be some classes whose flexible thresholds reach or approach , thereby filtering out most unlabeled data in the batch. The threshold warm-up solves this issue by gradually raising the thresholds of all classes from zero – it creates a learning boom at the early training stage where most of the unlabeled data can be utilized.
Comparison with class balancing objectives.
CPL has the effect of balancing across classes the number of unlabeled samples used to compute pseudo-labeling loss in each batch. Similar effect can be achieved by making the marginal class distribution close to a uniform distribution for each batch. We conduct such a comparative experiment by directly adding an additional objective to FixMatch: , where is the mean predicted probability of class across all samples in the batch, and is a uniform distribution: . The error rate of adding such an objective is 7.16% on the CIFAR-10 40-label split (compared with FixMatch 7.47%0.28 and FlexMatch 4.97%0.06). While this approach requires instances of each class within each batch to be balanced to make sense, CPL does not have such a constraint. It is more flexible and involves less human intervention to adjust thresholds than adjusting model’s predictions.
Related Work
Pseudo-Labeling is a pioneer SSL method that uses hard artificial labels converted from model predictions. A confidence-based strategy was used in along with pseudo labeling so that the unlabeled data are used only when the predictions are sufficiently confident. Such confidence-based thresholding also presents in recently proposed UDA and FixMatch with the difference being that UDA used sharpened ‘soft’ pseudo labels with a temperature whereas Fixmatch adopted one-hot ‘hard’ labels. The success of UDA and FixMatch, however, relies heavily on the usage of strong data augmentations to improve the consistency regularization. ReMixMatch also leveraged such strong augmentations.
The combination of curriculum learning and semi-supervised learning is popular in recent years . For multi-model image classification task, optimized the learning process of unlabeled images by judging their reliability and discriminability. In , the easy image-level properties are learned first and then used to facilitate segmentation via constrained CNNs. Curriculum learning is also used to alleviate out-of-distribution problems by picking up in-distribution samples from unlabeled data according to the out-of-distribution scores .
Several researches have investigated on dynamic threshold in related fields such as sentiment analysis and semantic segmentation . In , the threshold was gradually reduced to make high-quality data selected into labeled data set in the early stage and large-quantity in the later stage. An extra classifier is added to automate the threshold to deal with domain inconsistency in . introduced curriculum learning to self-training with a steadily increasing threshold and achieved near state-of-the-art results.
Conclusion and Future Work
In this paper, we introduce Curriculum Pseudo Labeling (CPL), a curriculum learning approach of leveraging unlabeled data for SSL. CPL dramatically improves the performance and convergence speed of SSL algorithms that involve thresholds while being extremely simple and almost cost-free. FlexMatch, our improved algorithm of FixMatch, achieves state-of-the-art performance on a variety of SSL benchmarks. In future work, we would like to improve our method under the long-tail scenario where the unlabeled data belonging to each class are extremely unbalanced.
Broader Impact
CPL fills the gap that no modern SSL algorithm considers the inherent learning difficulties of different classes during the training, and shows that by doing so, the convergence speed and final accuracy can both be improved. We hope that CPL can attract more future attention to explore the effectiveness of utilizing unlabeled data according to the model’s learning status as well as the per-class learning difficulty.
Funding Disclosure
Funding in direct support of this work: computing resource granted by Tokyo Institute of Technology and Microsoft Research Asia. This work was partially supported by Toray Science Foundation.
References
Appendix A Experimental Results
For reproduction, we show the detailed hyperparameter setting for each method in Table 3 and 4, for algorithm-dependent and algorithm-independent hyperparameters, respectively.
A.2 Class-wise accuracy improvement.
As introduced in the paper, CPL has its ability of improving performance on those hard-to-learn classes by taking into consider the model’s learning status. A detailed class-wise accuracy comparison is listed in Table 5, where the final accuracies of class 2, 3 and 5 with originally bad performance are improved.
A.3 Median error rates
We also report the median error rates of the last 20 checkpoints by allowing all methods to run the same iterations, following existing work . There are 1000 iterations between every two checkpoints. The results in Table 6 show that our CPL method can dramatically improve the performance of existing SSL algorithms and the FlexMatch achieves the best accuracy. These conclusions are in consistency with the results of Table1 in the main text, showing the effectiveness of our proposed CPL algorithm.
A.4 Detailed results
To comprehensively evaluate the performance of all methods in a classification setting, we further report the precision, recall, f1 score and AUC (area under curve) results on CIFAR-10 dataset. As shown in Table 7, we see that in addition to the reduced error rates, CPL also has the best performance on precision, recall, F1 score, and AUC. These metrics, together with error rates (accuracy), shows the strong performance of our proposed method.
Appendix B TorchSSL: A PyTorch-based SSL Codebase
The PyTorch framework has gained increasing attention in the deep learning research community. However, the main existing SSL codebase is based on TensorFlow. For the convenience and customizability, we re-implement and open source a PyTorch-based SSL toolbox, named TorchSSL Our toolbox is partially based on . as shown in Figure 5. TorchSSL contains eight popular semi-supervised learning methods: -Model , Pseudo-Labeling , VAT , Mean Teacher , MixMatch , ReMixMatch , UDA , and FixMatch , along with our proposed method FlexMatch. Most of our implementation details are based on . More importantly, in addition to the basic SSL methods and components, we implement several techniques to make the results stable under PyTorch framework. For instance, we add synchronized batch normalization to avoid the performance degradation caused by multi-GPU training with small batch size, and a batch norm controller to prevent performance crashes for some algorithms, which is not officially supported in PyTorch.
We observed that Mean Teacher can be very unstable if we update BatchNorm for both labeled data and unlabeled data in turn. Other algorithms such as -Model and MixMatch also show the similar instability. Therefore, we use BatchNorm Controller to update BatchNorm only for labeled data if labeled data and unlabeled data are forwarded separately. The code of BatchNorm Controller is as follows. We record the BatchNorm statistics before the forward propagation of unlabeled data and restore them after the propagation is done.
B.2 Benchmark results
We comprehensively run all algorithms in our TorchSSL on four common datasets in SSL: CIFAR-10, CIFAR-100, SVHN, and STL-10, and report the best error rates in Table 8, 9, 10, and 11, respectively. These benchmark results provide a reference of using this toolbox.