Guided Collaborative Training for Pixel-wise Semi-Supervised Learning
Zhanghan Ke, Di Qiu, Kaican Li, Qiong Yan, Rynson W. H. Lau
Introduction
Deep learning has been remarkably successful in many vision tasks. Nonetheless, collecting a large amount of labeled data for training is costly, especially for pixel-wise tasks that require a precise label for each pixel, e.g., the category mask in semantic segmentation and the clean picture in image denoising. Recently, semi-supervised learning (SSL) has become an important research direction to alleviate the lack of labels, by appending unlabeled data for training. Many SSL methods have been proposed for image classification with impressive results, including adversarial-based methods , consistent-based methods , and methods that are combined with self-supervised learning . In contrast, only a few works have applied SSL to specific pixel-wise tasks , and they mainly focus on semantic segmentation.
In this work, we investigate the generalization of SSL to diverse pixel-wise tasks. Such generalization is important in order for SSL to be used in new vision tasks with minimal efforts. However, generalizing existing pixel-wise SSL methods is not straightforward since they are designed for certain tasks by using task-specific properties (Sec. 2.2), e.g., assuming similar semantic contents between the input and output. Another possible generalization approach is to apply SSL methods designed for image classification to pixel-wise tasks. But there are two critical issues caused by the dense outputs, as illustrated in Fig. 1, leading to unsatisfactory performances of these methods on pixel-wise tasks.
First, dense outputs require pixel-wise prediction confidences (Sec. 2.3), which are difficult to estimate. Pixel-wise tasks are either pixel-wise classification (e.g., semantic segmentation and shadow detection) or pixel-wise regression (e.g., image denoising and matting). Although we may use the maximum classification probability to represent the prediction confidence in pixel-wise classification, it is unavailable in pixel-wise regression. Second, existing perturbations designed for SSL (Sec. 2.4) are not suitable for dense outputs. In pixel-wise tasks, strong perturbations in the input, e.g., clipping in Mean Teacher , will change the input image and its labels. As a result, the perturbed inputs from the same original image have different labels, which is undesirable in SSL. Besides, the perturbations through Dropout are disabled in most pixel-wise tasks. Although Dual Student proposes to create perturbations through different model initializations, its training strategy can only be used in image classification.
To address the above two issues caused by dense outputs, we propose a new SSL framework, named Guided Collaborative Training (GCT), for pixel-wise tasks. It includes three modules – two models for the specific task (the task models) and a novel flaw detector. GCT overcomes the two issues by: (1) approximating the pixel-wise prediction confidence by the output of the flaw detector, i.e., a flaw probability map, and (2) extending the perturbations used in Dual Student to pixel-wise tasks. Since different model initializations lead to inconsistent predictions for the same input, we can ensemble the reliable pixels, i.e., the pixels with lower flaw probabilities, in the predictions. In addition, minimizing the flaw probability map should help correct the unreliable pixels in the predictions. Motivated by these ideas, we introduce two SSL constraints, a dynamic consistency constraint between the task models and a flaw correction constraint between the flaw detector and each of the task models, to allow the modules in GCT to learn from unlabeled data collaboratively under the guidance of the flaw probability map rather than the task-specific properties. As a result, GCT can be applied to diverse pixel-wise tasks, simply by replacing the task models without structural adaptations.
We evaluate GCT on the standard benchmarks for semantic segmentation (pixel-wise classification) and real image denoising (pixel-wise regression). We also conduct experiments on our own practical datasets, i.e., the datasets with a large proportion of unlabeled data, for portrait image matting and night image enhancement (both are pixel-wise regression) to demonstrate the generalization of GCT on real applications. GCT surpasses start-of-the-art SSL methods that can be applied to these four challenging pixel-wise tasks. We envision that this work will contribute to future research and development of new vision tasks with scarce labels.
Related Work
Our work is related to two main branches of SSL methods designed for image classification. The adversarial-based methods assemble the discriminator from GAN , and try to match the latent distributions between labeled and unlabeled data through the image-level adversarial constraint. The consistent-based methods learn from unlabeled data by applying a consistency constraint to the predictions under different perturbations. Apart from them, some latest works combine self-supervised learning with SSL or expand the training set by interpolating labeled and unlabeled data .
2 SSL for Pixel-wise Tasks
Existing research on pixel-wise SSL mainly focuses on semantic segmentation. GANs dominate in this topic through the combination with the SSL methods derived from image classification. For example, Hung et al. extract reliable predictions to generate pseudo labels for training. Mittal et al. modify Mean Teacher to a multi-label classifier and use it as a filter to remove uncertain categories. Besides, Lee et al. and Huang et al. study weak-supervised learning in the SSL context. However, these works require pre-defined categories, which is a general property of classification-based tasks. Chen et al. apply SSL in face sketch synthesis, which belongs to pixel-wise regression. It regards the pre-trained VGG network as a feature extractor to impose a perceptual constraint on the unlabeled data. Unfortunately, the perceptual constraint can only be used in tasks that have similar semantic contents between the inputs and outputs. For example, it does not work on segmentation since the semantic content of the category mask is different from the input image.
3 Prediction Confidence in SSL
Prediction confidence is necessary for computing the SSL constraints, which consider the predictions with higher confidence values as the targets, i.e., pseudo labels. Earlier works show that the averaged targets are more confident. For example, Temporal Model accumulates the predictions over epochs as the targets; Mean Teacher defines an explicit model by exponential moving average to generate the targets; FastSWA further averages the models between epochs to produce better targets. Others regard the maximum classification probability as the prediction confidence.
In pixel-wise SSL, the outputs of the discriminator are used to approximate the prediction confidence . Instead, we propose the flaw detector to estimate the prediction confidence, with two key differences. First, the flaw detector predicts a dense probability map with location information while the discriminator predicts an image-level probability. Second, we use the ground truth of the labeled data to generate the targets of the flaw detector.
4 Perturbations in SSL
Many SSL methods heavily rely on perturbations for training. The consistent-based methods utilize data augmentations to alter the inputs. To further improve the inconsistency, VAT generates virtual adversarial noises while S4L adds a rotation operation to the inputs. Others such as MixMatch and ReMixMatch generate perturbed samples by data interpolation. Apart from the perturbations in the inputs, Dropout perturbs the predictions through a random selection of nodes . The models in Dual Student have inconsistent predictions for the same input due to different initializations.
Since the perturbations from both data augmentations and Dropout are not suitable for dense outputs, GCT follows Dual Student in creating perturbations. However, unlike Dual Student, GCT learns from unlabeled data through the two SSL constraints based on the flaw detector, allowing GCT to be applicable to diverse pixel-wise tasks.
Guided Collaborative Training
In this section, we first present an overview of GCT. We then introduce the flaw detector and the two proposed SSL constraints. Fig. 2 shows the GCT framework. and are the two task models, which are referred to as () in the following context. The architecture of is arbitrary, and GCT allows the task models to have different architectures. The only requirement is that and should have different initializations to form the perturbations between them (which is the same as Dual Student). denotes the flaw detector. In SSL, we have a dataset consisting of a labeled subset with labels and an unlabeled subset . The inputs for both and are exactly the same. Given an , the GCT framework first predicts of size , where the value of is defined by the specific task. Then, the concatenation of and is processed by to estimate the flaw probability map of size . The prediction confidence map can be approximated by . We train GCT iteratively in two steps like GAN .
In the first step, we train with fixed . For the labeled data, the prediction is supervised by its corresponding label as:
where is a task-specific constraint, and is a pixel index. To learn the unlabeled data, we propose a dynamic consistency constraint and a flaw correction constraint , which are guided by the flaw probability map and will be described in Sec. 3.3 and Sec. 3.4, respectively. The final constraint for is a combination of three constraints as:
where is a pair of labeled data. and are hyper-parameters to balance the two SSL constraints.
In the second step, learns from the labeled subset. We calculate the ground truth of through a classical image processing pipeline based on and . In our framework, is trained by using Mean Square Error (MSE) as:
where is the ground truth of , which will be discussed in Sec. 3.2.
2 Flaw Detector
On the labeled subset, the goal of the flaw detector is to learn the flaw probability map that indicates the difference between and , i.e., the flaw regions in . One simple way to find the flaw regions is . However, it is difficult to learn many tasks since it is sparse and sharp (column (4) of Fig. 9). To address this problem, we introduce an image processing pipeline that converts to a dense probability map (column (5) of Fig. 9). consists of three basic image processing operations: dilation, blurring and normalizationRefer to Appendix A in the Supplementary for the algorithm of .. To estimate the flaw probability map for the unlabeled data, we apply a common SSL assumption : the distribution of unlabeled data is the same as that of the labeled data. Therefore, trained on the labeled subset should also work well on the unlabeled subset.
The architecture Refer to Appendix B in the Supplementary for the architecture of the flaw detector. of the flaw detector is similar to the fully convolutional discriminator in . However, averages all predicted pixels to get a single confidence value during training, as its target is an image-level real or fake probability. In pixel-wise tasks, the prediction is usually accurate for some pixels but not the others, and pixels of higher accuracy should have higher confidence. Using an average confidence to represent the overall confidence is not appropriate. For example, may be more confident (more accurate) than in a small local region although the average prediction confidence of is lower than . Therefore, the per-pixel prediction confidence (from the flaw detector) is more meaningful than the average prediction confidence (from the discriminator) in pixel-wise tasks. Fig. 9 visualizes the results of and in the four validated tasks.
3 Dynamic Consistency Constraint
The two task models in GCT have inconsistent predictions for the same input due to the perturbations between them. We use the dynamic consistency constraint to ensemble the reliable pixels in and . Typically, the standard consistency constraint is unidirectional, e.g., from the ensemble model to the temporary model. Here, “dynamic” indicates that our is bidirectional and its direction changes with the flaw probability (Fig. 4(a)). Intuitively, if a pixel in has a lower flaw probability, we treat it as the pseudo label to the corresponding pixel in . To assure the quality of the pseudo label, we introduce a flaw threshold to disable for the pixels that have higher flaw probability values than in both and . Through this process, there is an effective knowledge exchange between the task models, making them collaborators.
Formally, given a sample , GCT outputs , , and their corresponding flaw probability maps , through forward propagation. We first normalize the values in to $\xi1$ as:
is a boolean-to-integer function, which outputs 1 when the is true and 0 otherwise. We define the dynamic consistency constraint for as:
4 Flaw Correction Constraint
Apart from , the flaw correction constraint attempts to correct the unreliable predictions of the task models (Fig. 4(b)). The key idea behinds is to force the values in the flaw probability map to become zero. We define for (with being fixed) as:
We use a binary mask to enable on the pixels without , i.e., the pixels with unreliable predictions in both task models:
We consider that the flaw detector helps improve the task models through . For a system containing only one task model and the flaw detector, the objectives can be derived from Eq. (3) and (6) as:
where and have the same distribution. We simplify Eq. (8) by removing the pixel summation operation. In such situation, learns the flaw probability map while optimizes it with a zero label. If we assume that the training process converges to an optimal solution in iteration , we have:
where is the current iteration. Hence, the objective changes during the training process and is equal to when . The alignment in the objectives indicates that and are collaborative to some degree.
To illustrate the difference between and the adversarial constraint, we compare Eq. (8) with the objectives of LSGAN . If we modify LSGAN for SSL, its objectives should be:
where is the standard discriminator that tries to differentiate and . In contrast, tries to match the distributions between and . Here we reverse the labels, i.e., 1 for fake and 0 for real, to be consistent with Eq. (8). Since the targets of are constants, we have:
which means that and are adversarial during the whole training process.
Experiments
In order to evaluate our framework under different ratios of the labeled data, we experiment on the standard benchmarks for semantic segmentation and real image denoising. We also experiment on the practical datasets created for portrait image matting and night image enhancement to demonstrate the generalization of GCT in real applications. We further conduct ablation experiments to analyze various aspects of GCT.
Experimental Setup. We notice that existing works of pixel-wise SSL usually report a fully supervised baseline with a lower performance than the original paper due to inconsistent hyper-parameters. In image classification, a similar situation has been discussed by . To fairly evaluate the performance of SSL, we define some training rules to improve the SupOnly baselines. We denote the total number of trained samples as , where is the training epochs, is the number of iterations in each epoch, and is the batch size, which is fixed in each task. For the experiments performed on the standard benchmarks:
We train the fully supervised baseline according to the hyper-parameters from the original paper to achieve a comparable result. The same hyper-parameters (except ) are used in (2) and (3).
We use the same as in (1) to train the models supervised by the labeled subset (SupOnly). Although decreases as the labeled data reduces, to prevent overfitting, we do not increase by training more epochs.
We adjust to ensure that in SSL experiments is the same as (1). In SSL experiments, each batch contains both labeled and unlabeled data. We define “epoch” as going through the unlabeled subset for once. Meanwhile, the labeled subset is repeated several times inside an epoch.
By following these rules, the SupOnly baselines obtain good enough performance and do not overfit. The models trained by SSL methods have the same computational overhead, i.e., the same , as the fully supervised baseline. For experiments on the practical datasets, we first train epochs for the SupOnly baselines. Afterwards, we train the SSL models with the same . We use the grid search to find suitable hyper-parameters for all SSL methods. Refer to Appendix C in the Supplementary for more training details.Refer to Appendix D in the Supplementary for visual comparisons..
Semantic segmentation takes an image as input and predicts a series of category masks, which link each pixel in the input image to a class (Fig. 9(a)). We conduct experiments on the Pascal VOC 2012 dataset , which comprises 20 foreground classes along with 1 background class. The extra annotation set from the Segmentation Boundaries Dataset (SBD) is combined to expand the dataset. Therefore, we have 10,582 training samples and 1,449 validation samples. During training, the input images are cropped to after random scaling and horizontal flipping. Following previous works , we use DeepLab-v2 with the ResNet-101 backbone as the SupOnly baselines and as the task model in SSL methods. The same configurations as the original paper of DeepLab-v2 are applied, except the multi-scale fusion trick.
For SSL, we randomly extract , , , samples as the labeled subset, and use the rest of the training set as the unlabeled subset. Note that the same data splits are used in all SSL methods. Table 1 shows the mean Intersection-over-Union (mIOU) on the PASCAL VOC 2012 dataset with pre-training on the Microsoft COCO dataset . GCT achieves a performance increase of (under labels) to (under labels) over the SupOnly baselines. Moreover, our fully supervised baseline () is comparable with the original paper of DeepLab-v2 (), which is better than the result reported in (). Therefore, all SSL methods only have slight improvement under the full labels.
2 Real Image Denoising Experiments
Real image denoising is a task that devotes to removing the real noise, rather synthetic noise, from an input natural image (Fig. 9(b)). We conduct experiments on the SIDD dataset , which is one of the largest benchmarks on real image denoising. It contains 160 image pairs (noisy image and clean image) for training and 40 image pairs for validation. We split each image pair into multiple patches with size for training. The total training samples is about 30,000. We use DHDN , a method that won the second place in the NTRIE 2019 real image denoising challenge , as the task model since the code for the first place winner has not been published. The peak-signal-to-noise-ratio (PSNR) is used as the validation metric.
In image denoising, even small errors between the prediction and the ground truth can result in obvious visual artifacts. It means that the reliable pseudo labels are difficult to obtain, i.e., this task is difficult for SSL. We notice that the task models with the same architecture in GCT have similar predictions. Therefore, the perturbations from different initializations are not strong enough. To alleviate this problem, we replace one of the task models with DIDN that won the third place in the NTRIE 2019 challenge. We still use DHDN for validation.
We extract , , , labeled image pairs randomly for SSL. As shown in Table 2, our fully supervised baseline achieves 39.38dB (PSNR), which is comparable with the top-level results on the SIDD benchmark. Although SSL shows limited performance in this difficult task, GCT surpasses other SSL methods under all labeled ratios. Notably, GCT improves on PSNR by 0.61dB with labels (only 10 labeled image pairs) while the previous SSL methods improve on PSNR by 0.33dB at most.
3 Portrait Image Matting Experiments
Image Matting predicts a foreground mask (matte) from an input image and a pre-defined trimap. Each pixel value in the matte is a probability between $x\sim$20min per image). Finally, we combine 100 labeled images with 7,700 unlabeled images as the training set, while the remaining 200 labeled images are used as the validation set. For each labeled image, we generate 15 samples by random cropping and 35 samples by background replacement (with the OpenImage dataset ). For each unlabeled image, we generate 5 samples by random cropping. The structure of our task model is derived from , which is a milestone in image matting.
In this task, we verify the impact of increasing the amount of unlabeled data on SSL by experimenting on two configurations. With 100 labeled images, (1) we randomly select half (3,850) of unlabeled images for training, and (2) we use all (7,700) unlabeled images for training. As shown in Table 3, GCT yields an improvement over the SupOnly baselines by 1.96dB and 3.99dB for 3,850 and 7,700 unlabeled images respectively. This indicates that the SSL performance can be effectively improved by increasing the amount of unlabeled data. In addition, doubling the amount of unlabeled images achieves a more significant improvement (2.03dB) with GCT, compared with existing SSL methods.
4 Night Image Enhancement Experiments
Night Image Enhancement is another common vision application. This task adjusts the coefficients of the channels in a night image to show more details (Fig. 9(d)). Our dataset contains 1,900 night images captured by smartphones, of which 400 images are labeled using Photoshop (15min per image). We combine 200 labeled images with 1,500 unlabeled images for training and use another 200 labeled images for testing. We use horizontal flipping, slight rotation, and random cropping (to ) as data augmentations during training. We regard HDRNet as the task model. Since the dataset is small, we experimented with only one SSL configuration (Table 3). Similar to the experiments in the other three tasks, GCT outperforms existing SSL methods.
5 Ablation Experiments
We conduct ablation studies to analyze the proposed SSL constraints, the hyper-parameters in GCT, and the combination of the flaw detector and Mean Teacher.
Effect of the SSL Constraints. By default, GCT learns from the unlabeled data through the two SSL constraints simultaneously. In Fig. 5, we compare the experiments of training GCT with only one SSL constraint on the benchmarks for semantic segmentation and real image denoising. The results demonstrate that both and are effective. GCT with boosts the performance impressively, proving that the knowledge exchange between the two task models is reliable and effective. Meanwhile, the curve of GCT with indicates that the flaw detector also plays a vital role in learning the unlabeled data. Moreover, combining and allows GCT to achieve the optimal performance.
Hyper-parameters in GCT. We analyze the two hyper-parameters required by GCT (mentioned in Sec. 3.3), the flaw threshold and the cosine ramp-up epochs of , on the Pascal VOC benchmark for semantic segmentation with labels. Table 4 (left) shows the results under different , which controls the combination of the two SSL constraints. Specifically, only is applied when , and only is applied when . Our experiments show that can be set roughly, e.g., is suitable for semantic segmentation. The cosine ramp-up with epochs prevents exchanging unreliable knowledge due to the non-convergent flaw detector in the early training stage. The results in Table 4 (right) indicate that GCT is robust to , even though the cosine ramp-up is necessary for the best performance.
Combination of the Flaw Detector and MT. The consistency constraint in MT is applied from the teacher model to the student model. However, the teacher model may be worse than the student model on some pixels, which may cause a performance degradation. To avoid this problem, we use the flaw detector to disable the consistency constraint when the flaw probability of the teacher’s prediction is larger than the student’s prediction. Under labels, this method improves the mIOU value of MT from to on Pascal VOC and improves the PSNR value of MT from 38.22dB to 38.42dB on SIDD.
Conclusions
We have studied the generalization of SSL to diverse pixel-wise tasks and indicated the drawbacks of existing SSL methods in these tasks, which to the best of our knowledge is the first. We have presented a new general framework, named GCT, for pixel-wise SSL. Our experiments have proved its effectiveness in a variety of vision tasks. Meanwhile, we also note that SSL still has limited performance for tasks that require highly precise pseudo labels, such as image denoising. A possible future work is to investigate this problem and explore ways to create more accurate pseudo labels.
Guided Collaborative Training for Pixel-wise Semi-Supervised Learning Supplementary Material
Zhanghan Ke Di Qiu Kaican Li Qiong Yan Rynson W.H. Lau
Appendix A: Algorithm of C𝐶C
In GCT, we use a classical image processing pipeline to calculate the ground truth of the flaw detector on the labeled subset by taking the task model prediction and the corresponding label as the input. is composed of three operations:
: Blur by a Gaussian kernel of given shape.
: Dilate for each local region of given shape.
: Normalize all pixels in to range between .
We show the pseudo code of in Python style as follows (assume the shape of is ):
In our experiments, we set for semantic segmentation, and we set for other three tasks. We set for real image denoising, for night image enhancement, and for other two tasks.
Appendix B: Architecture of Flaw Detector
The flaw detector is a fully-convolutional neural network, which contains 8 convolutional layers with kernels. The amount of kernels is increased from to in the first layers and then decreased to in the last layer. Each of the first 7 convolutional layers is followed by batch normalization and leaky ReLU with threshold of . The convolutional layers with stride= reduce the resolution of the feature maps. At the end of , we add a bilinear interpolation operation to rescale the output to the size of the input. In all experiments of GCT, we optimize by Adam (with learning rate ). The architecture of is as follow:
Appendix C: Training Details
We have experimented with several SSL methods, including (1) the consistent-based Mean Teacher (MT) ; (2) the self-supervised SSL (S4L) ; (3) the adversarial-based method proposed in (AdvSSL); (4) the GCT framework proposed by us. Here are the definitions of the hyper-parameters for SSL in these methods:
For the four validated tasks, we use grid search to find the suitable hyper-parameters for SSL. The final settings for the experiments are as follows:
Appendix D: Visual Comparisons
Here we provide visual comparisons of the SSL results for four validated tasks. The red bounding box in the figure highlights some main differences in the outputs. As shown below, GCT surpasses existing SSL methods in visual effects.