Self-supervised Pre-training with Hard Examples Improves Visual Representations
Chunyuan Li, Xiujun Li, Lei Zhang, Baolin Peng, Mingyuan Zhou, Jianfeng Gao
Introduction
Self-supervised visual representation learning aims to learn image features from raw pixels without relying on manual supervisions. Recent results show that self-supervised pre-training (SSP) outperforms state-of-the-art (SoTA) fully-supervised pre-training methods , and is becoming the building block in many computer vision applications. The pre-trained model produces general-purpose features and serve as the backbone of various downstream tasks such as classification, detection and segmentation, improving the generalization of those task-specific models that are often trained on limited amounts of task labels.
Most state-of-the-art SSP methods focus on designing novel pretext objectives, ranging from the traditional prototype learning , to a recently popular concept known as contrastive learning , and a combination of both . Apart from the improved efficiency, all these methods heavily rely on data augmentation to create different views of an image using image transformations, such as random crop (with flip and resize), color distortion, and Gaussian blur. Recent studies show that SSP performance can be further improved by using more aggressively transformed views, such as increasing the number of views , and more distinctive views via minimizing mutual information .
However, image transformations are agnostic to the pretext objectives, and it remains unknown how to augment views specifically based on the pre-training tasks themselves, and how different augmentation methods affect the generalization of the learned models. To tailor data augmentation to pre-training tasks, we explicitly formulate SSP as a problem of predicting pseudo-labels, based on which we propose to generate hard examples (Hexa), a family of augmented views whose pseudo-labels are difficult to predict. Specifically, two schemes are considered. Adversarial examples are created with the intention to cause an SSP model to make prediction mistakes and thus improve the generalization of the model . Cut-mixed examples are created via cutting and pasting patches among different images , so that its content is a mixture of multiple images.
Our contributions include: A pseudo-label perspective is formulated to motivate the concept of hard examples in self-supervised learning. Two novel algorithms are proposed, through applying our framework to two distinctly different existing approaches. Experiment are conducted on a wide range of tasks in self-supervised benchmarks, showing that Hexa consistently improves their original counterparts, and achieves SoTA performance under the same settings. It demonstrates the genericity and effectiveness of proposed framework in constructing hard examples for improving the visual representations using SSP.
Self-supervision: A Pseudo-label View
Contrastive learning is a framework that learns representations by maximizing agreement between differently augmented views of the same image via a contrastive loss in the latent space. For a given query , we identify its positive samples from a set of keys , where positive samples are indexed by and negative samples are indexed by . The pseudo-labels in contrastive learning are defined by feature pairwise comparisons: for the pair and for the pair . For a query with pairs, its pseudo-label vector is .
where denotes the set of trainable parameters and is a temperature hyper-parameter. From (1) to (2), only the loss term indexed with remains, while the ones indexed with are excluded, because their corresponding pseudo-label .
In the instance discrimination pretext task (used by MoCo and SimCLR), a query and a key form a positive pair if they are data-augmented versions of the same image, and otherwise form a negative pair. The contrastive loss (2) can be minimized by various mechanisms that differ in how the keys (or negative samples) are maintained .
SimCLR The negative keys are from the same batch and updated end-to-end by back-propagation. SimCLR is based on this mechanism and requires a large batch to provide a large set of negatives.
MoCo In the MoCo mechanism, the negative keys are maintained in a queue , and only the queries and positive keys are encoded in each training batch. A momentum encoder is adopted to improve the representation consistency between the current and earlier keys. MoCo decouples the batch size from the number of negatives. MoCo-v2 is an improved version using strong augmentation (i.e. more aggressive image transformations) and MLP projection proposed in SimCLR.
Type II: Prototype Learning.
The prototype learning methods introduce a “prototype” as the centroid for a cluster formed by similar image views. The latent representations are fed into a clustering algorithm to produce the prototype/cluster assignments, which are subsequently used as “pseudo-labels” to supervise representation learning.
DeepCluster is a representative prototype learning work. It employs -means as the clustering algorithm, which takes a set of latent vectors as input, clusters them into distinct groups with prototypes , and simultaneously output the optimal cluster assignment as a one-hot probability simplex. The model is trained to predict the optimal assignment:
Pre-training with Hard Examples
where is a hyper-parameter governing how invariant the resulting model should be to adversarial attacks, and is the perturbation. In practice, (5) is updated using two steps: By applying Projected Gradient Descent (PGD) , we obtain adversarial examples on-the-fly:
Adversarial Prototype Learning.
The adversarial training for prototype-based methods are similar to supervised settings, after the cluster assignments are learned. We treat these pseudo-labels as targets to fool the model:
Implementations.
2 Cut-Mixed Examples
3 Full Hexa Objective
The overall self-adversarial training objective considers both clean and hard examples constructed by adversarial and cutmix augmentations:
where and are the weighting hyper-parameters to control the effect of adversarial examples and cutmixed examples, respectively. In our experiments, we set and/or . Note that reduces the objective to the standard self-supervised training algorithms. Concretely, we consider two novel algorithms:
Hexa By plugging terms (2) (8) and (5) into (10), it yields the full self-adversarial contrastive learning objective denoted as . The Hexa training procedure is detailed in Algorithm 1. We build Hexa on top of MoCo-v2. The two algorithms are distinguished from each other in Lines 5-9, where hard examples are computed on query and subsequently employed in model update for Hexa.
Hexa The full self-adversarial prototype learning objective is obtained via plugging (4)(8) and (7) into (10). We build Hexa based on DeepCluster-v2 , which improves DeepCluster to reach similar performance with recent state-of-the-art methods. The Hexa training procedure is detailed in Algorithm 2. It differs from DeepCluster-v2 in Lines 6-10, where hard examples are computed to train the network in conjunction with clean examples.
Related Works
Self-supervised learning is a popular form of unsupervised learning, where labels annotated by humans are replaced by “pseudo-labels” directly extracted from the raw input data by leveraging its intrinsic structures. We broadly categorize existing self-supervised learning methods into three classes: Handcrafted pretext tasks. This includes many traditional self-supervised methods such as relative position of patches , masked pixel/patch prediction , auto-regressive modeling , rotation prediction , image colorization , cross-channel prediction and generative modeling . These approaches typically exploit domain knowledge to carefully design a pretext task, with the learned features often focusing on one certain aspect of images, leading to a limited transfer ability. Contrastive learning. The instance-level classification task is considered , where each image in a dataset is treated as a unique class, and various augmented views of an image are the examples to be classified. Some recent works in this line are CPC , deep InfoMax , MoCo , SimCLR , BYOL etc. Prototype learning. Clustering is employed for deep representation learning, including DeepCluster , SwAV and PCL , among many others . The proposed Hexa can be generally applied to all three classes in principle, as long as the notation of pseudo-labels exists. In this paper, we focus on the latter two classes, as they have shown SoTA representation learning performance, surpassing the ImageNet-supervised counterpart in multiple downstream vision tasks.
The role of augmentations.
Image data augmentations/transformations such as crop and blurring play a crucial role in modern self-supervised learning pipeline. It has been empirically shown that visual representations can be improved by employing stronger image transformations and increasing the number of augmented views of an image . InfoMin studied the principles of good views for contrastive learning, and suggested to select views with less mutual information. By definition, adversarial and cut-mixed examples tend to be harder examples than transformation-augmented ones for self-supervised problems, and are complementary to the above techniques.
2 Hard Examples
A vast majority of works commonly view adversarial examples as a threat to models , and suggest training with adversarial examples leads to accuracy drop on clean data . Adversarial training have been studied for self-supervised pre-training . Our work is significantly different from Chen et al. in two aspects: Motivations – We aim to use adversarial examples to boost standard recognition accuracy on large-scale datasets such as ImageNet, while Chen et al. mainly study model robustness on small datasets such as CIFAR-10. Algorithms – We focus on the modern contrastive/prototype learning methods (last two categories of SSP methods in Section 4.1), while Chen et al. work on traditional handcrafted SSP methods (the first category).
Improved standard accuracy.
Hard examples have been shown to be effective in improving recognition accuracy in supervised learning settings. For adversarial examples, one early attempt is virtual adversarial training (VAT) , a regularization method that improves semi-supervised learning tasks. The success was recently extended to natural language processing and vision-and-language tasks . In computer vision, AdvProp is a recent work showing that adversarial examples improve recognition accuracy on ImageNet in supervised settings. Hadi et al. further show that adversarially robust ImageNet models transfer better . For cut-mixed examples, it was first studied by Yun et al. . Similar augmentation schemes using a mixture of images include mix-up , cut-out etc. All above hard examples are constructed in the supervised settings, our Hexa is the first work to systematically study hard examples in large-scale self-supervised settings, due to the proposed pseudo-label formulation. We confirm that hard examples improve the model’s transfer ability.
Experimental Results
All of our study for unsupervised pretraining (learning encoder network without labels) is done using the ImageNet ILSVRC-2012 dataset . We implement Hexa based on the pre-training scheldule of MoCo-v2, and implement Hexa based on the pre-training scheldule of DeepCluster-v2. Both use the cosine learning rate and MLP projection head. Due to the limit of computational resource, all experiments are conducted with ResNet-50 and pre-trained in 200/800 epochs if not specifically mentioned. Once the model is pre-trained, we follow the same fine-tuning protocols/schedules with the baseline methods . Following common practice in evaluating pre-trained visual representations, we test the model’s transfer learning ability on a wide range of datasets/tasks in the self-supervised learning benchmark , based on the principle that a good representation should transfer with limited supervision and limited fine-tuning.
In what follows, we denote Hexa and Hexa as two variants that are both constructed with 2 random crops at resolution 224 and adversarial examples. More specifically, Hexa follows MoCo-v2: one crop for query and the other for key; Hexa is always compared with the DeepCluster-v2 variant with 2 crops. The current SoTA method is SwAV , which employs 8 random crops: 2 crops at resolution 224 and 6 crops at resolution 96. To compare with SoTA, we also increase the number of crops to 8 and consider two variants: Hexa(8-crop) is with adversarial examples, and Hexa(8-crop) is constructed with both adversarial and cut-mixed examples. All 2-crop methods use a mini-batch size of B=256, and 8-crop methods use a mini-batch size of B=4096.
2 Linear classification
To evaluate the learned representations, we first follow the widely used linear evaluation protocol, where a linear classifier is trained on top of the frozen base network, and test accuracy is used as a proxy for representations. We follow previous setup and evaluate the performance of such linear classifiers on four datasets, including ImageNet , PASCAL VOC2007 (VOC07) , CIFAR10 (C10) and CIFAR100 (C100) . A softmax classifier is trained for ImageNet/CIFAR, while a linear SVM is trained for VOC07. We report 1-crop (), Top-1 validation accuracy for ImageNet/CIFAR and mAP for VOC07.
Table 1 shows the results of linear classification. It is interesting to observe that DeepCluster-v2 is slightly better than MoCo-v2, indicating that the traditional prototype methods can be on par with the popular contrastive methods, with the same pre-training epochs and data augmentation strategies. We hope this result can inspire future research to more carefully select different pretext objectives. By contrast, Hexa variants consistently outperform theirs counterparts for both contrastive and prototype methods, demonstrating that the proposed hard examples can effectively improve learned visual representations in SSP.
We also pre-train Hexa with 800 epochs, a longer schedule used in MoCo-v2 . The learning curves are compared in Figure 3(a). Hexa is consistently better than MoCo-v2 and the gap is larger at the beginning. We hypothesize that the augmentation space is more efficiently explored with hard examples than with traditional image transformations, but this advantage is less reflected in improved recognition accuracy, when the augmentation space is gradually fully occupied at the end of training. When comparing with SoTA methods equipped with multi-crop , we see that Hexa(8-crop) achieves slightly better than SwAV on ImageNet, and even outperforms InfoMin with 800 pre-training steps. By plotting the training curves of their linear classifiers in Figure 3(b), we observe that Hexa(8-crop) clearly outperforms SwAV with limited fine-tuning (e.g. 20 epochs training). The advantage of Hexa is more significantly than SwAV with limit supervision, this can be seen from a larger performance gap on VOC07 in Table 1.
We evaluate the learned representation on image classification tasks with few training samples per-category. We follow the setup in Goyal et al. and train linear SVMs using fixed representations on VOC07 for object classification. We vary the number of training samples per-class and report the average result across 5 independent runs. The results are shown in Table 2. Hard examples help improve the performance for both contrastive and prototype learning, especially when . This is probably because the performance is very sensitive to the choice of selected labelled samples when , rendering the evaluation less stable. Pre-training longer (MoCo-v2 with 800 epochs) helps reduce this issue, and the proposed hard examples can further boost the performance. When samples are considered, the proposed scheme surpasses the ImageNet-supervised pre-training approach. To the best of our knowledge, Hexa is the first work to surpasses the supervised baseline with such a small number of labelled samples on VOC07, showing high sample-efficiency of the learned representations. Hexa pre-trained at 200 epochs also outperforms SwAV (pre-trained at both 200 epochs and 800 epochs) by a large margin in all cases.
3 Semi-supervised learning on ImageNet
We perform semi-supervised learning experiments to evaluate whether the learned representation can provide a good basis for fine-tuning. Following the setup from Chen et al. , we select a subset (1% or 10%) of ImageNet training data (the same labelled images with Chenet al. ), and fine-tune the entire self-supervised trained model on these subsets. For the proposed Hexa, and we fine-tune the models using the same schedule. SwAV with 8 augmentation crops and 200 pre-training epochs is used a fair baseline.
Table 3 reports the Top-1 and Top-5 accuracy on ImageNet validation set. Hexa improves its counterparts MoCo-v2 and DeepCluster-v2 in all cases. By different variants of Hexa, we see that cut-mixed examples are important in boosting performance, especially with 1% labels. Hexa(8-crop ) sets a new SoTA under 200 training epochs, outperforming all existing self-supervised learning methods. It even outperforms BYOL pre-trained at 800 epochs in both cases. For SwAV pre-trained at 200 epochs, it is significantly inferior to Hexa in the same setting. For SwAV pre-trained at 800 epochs, it achieves Top-1 53.9% and Top-5 78.5% with 1% labelled images, which is lower than our Hexa pre-trained at 200 epochs by a notable margin. This again shows the effectiveness of hard examples in improving visual representations in low-resource settings.
We also fine-tune over 100% of ImageNet labels for 20 epochs, and Hexa reaches 78.6% Top-1 accuracy, outperforming the supervised approach (76.5%) using the same ResNet-50 architecture by a large margin (2.1% absolute recognition accuracy). Hexa also achieves higher performance compared with all existing self-supervised learning methods in both 200 and 800 pre-training epochs settings. This shows that hard examples can effectively improve SSP, which can be viewed as a promising approach to further improve standard supervised learning such as Big Transfer in the future.
4 Object detection
It is standard practice in data-scarce object detection tasks to initialize earlier model layers with the weights from ImageNet-trained networks. We study the benefits of using hard-examples-trained networks to initialize object detection. On the VOC object detection task, a Faster R-CNN detector is fine-tuned end-to-end on the VOC 07+12 trainval set1 and evaluated on the VOC 07 test set using the COCO suite of metrics . The results are shown in Table 4. We find that Hexa consistently outperforms MoCo-v2 that is pre-trained with standard image transformations.
Conclusion
We have presented a comprehensive study of utilizing hard examples to improve visual representations for image self-supervised learning. By treating SSP as a pseudo-label classification task, we introduce a general framework to generate harder augmented views to boost the discriminative power of self-supervised learned models. Two novel algorithmic variants are proposed: Hexa for contrastive learning and Hexa for prototype learning. Our Hexa variants outperform their counterparts, often by a notable margin, and achieve SoTA under the same settings. Future research directions include incorporating more advanced hard examples under this framework, and exploring their performance with larger networks.
The authors gratefully acknowledge Bai Li for helpful discussion. Additional thanks go to the entire Project Philly team inside Microsoft, who provided us the computing platform for our research.
References
Appendix A Hyper-parameter Choice
We study the hyper-parameter choices attack perturbation threshold and PGD step size in adversarial images. For Hexa, we grid search over and . Each variant is pre-trained for 40 epochs. A linear classifier is added on the pre-trained checkpoint and trained for 100 epochs. The results are shown Figure 4. Adding too large perturbations can hurt model performance significantly. Otherwise, the model perform similarly with differnt ways of adding small perturbations, allowing either a large threshold with a small step size, or a large step size with a small threshold. We used for convenience.
Cut-mixed examples.
We study the hyper-parameter choices in in cut-mixed images. We consider 6 random crops for each image: 2 crops at resolution 160 and 4 crops at resolution 96. The model is pre-trained in 5 epoch, and a linear classification on the checkpoint is trained for 1 epoch. The results are shown in Figure 5. Cut-mixed examples in various settings improves performance. We used in our experiments.
Appendix B Experiments details for transfer learning
The main network is fixed, and global average pooling features (2048-D) of ResNet-50 are extracted. We train for 100 epochs. For Hexa, we follow the block-decay training schedule of with an initial learning rate of 30.0 and step decay with a factor of 0.1 at ${}_{\text{DCluster}}$, we follow the cosine-decay training schedule of with an initial learning rate of 0.3. The logistic regression classifier is trained using SGD with a momentum of 0.9.
Linear classification on VOC07
For training linear SVMs on VOC07, we follow the procedure in and use the LIBLINEAR package . We pre-process all images by resizing to 256 pixels along the shorter side and taking a center crop. The linear SVMs are trained on the global average pooling features of ResNet-50.
Linear classification on Cifar10 and Cifar100
We trained a linear classifier on features extracted from the frozen pre-trained network. We used Adamax to optimize the softmax cross-entropy objective for 20 epochs, a batch size of 256, a learning rate [0.1,0.01,0.001] and decay at with a factor of 0.1. All images were resized to pixels (after which we took a center crop), and we did not apply data augmentation.
Semi-supervised learning on ImageNet
We follow to finetune ResNet-50 with pretrained weights on a subset of ImageNet with labels. We optimize the model with SGD, using a batch size of 256, a momentum of 0.9, and a weight decay of 0.0005. We apply different learning rate to the ConvNet and the linear classifier. The learning rate for the ConvNet is 0.01, and the learning rate for the classifier is 0.1 (for 10% labels) or 1 (for 1% labels). We train for 20 epochs, and drop the learning rate by 0.2 at 12 and 16 epochs.
Object detection on VOC
We follow to use the R50-FPN backbone for the Faster R-CNN detector available in the Detectron2 codebase . We freeze all the conv layers and also fix the BatchNorm parameters. The model is optimized with SGD, using a batch size of 8, a momentum of 0.9, and a weight decay of 0.0001. The initial learning rate is set as 0.05. We finetune the models for 15 epochs, and drop the learning rate by 0.1 at 12 epochs.
Appendix C Experiments on Fine-tuning
We fine-tuned the entire network using the weights of the pre-trained network as initialization. We trained for 20 epochs at a batch size of 256 using Adamax, decayed at with a factor of 0.1. We grid search learning rate over [0.0005, 0.001, 0.01]. The results are shown in Table 5. Our Hexa consistently improves their original counterparts for both datasets.