Contrastive Learning with Adversarial Examples
Chih-Hui Ho, Nuno Vasconcelos
Introduction
Deep networks have enabled significant advances in many machine learning tasks over the last decade. However, this usually requires supervised learning, based on large and carefully curated datasets. Self-supervised learning (SSL) self_supervised_survey aims to alleviate this limitation, by leveraging unlabeled data to define surrogate tasks that can be used for network training. Early advances in SSL were mostly due to the introduction of many different surrogate tasks egomotion ; PathakKDDE16 ; Hsin17 ; UEL ; coloring ; LarssonMS16 , including solving image NorooziF16 ; damagedPuzzle or video VideoJigsaw puzzles, filling image patches PathakKDDE16 ; Doersch2015UnsupervisedVR ; Mundhenk17 or discriminating between image rotations rotation . Recently, there have also been advances in learning techniques specifically tailored to SSL, such as contrastive learning (CL) UEL ; simclr ; he2019moco ; Insdis ; tian2019contrastive ; TCL ; TCN ; misra2020pirl , which is the focus of this paper. CL is based on a surrogate task that treats instances as classes and aims to learn an invariant instance representation. This is implemented by generating a pair of examples per instance, and feeding them through an encoder, which is trained with a constrastive loss. This encourages the embeddings of pairs generated from the same instance, known as positive pairs, to be close together and embeddings originated from different instances, known as negative pairs, to be far apart.
The design of positive pairs is one of the research focuses of CL contrast_theory . These pairs can include data from one or two modalities. For single-modality approaches, a common procedure is to rely on data augmentation techniques. For example, instances from image datasets are frequently subject to transformations such as rotation, color jittering, or scaling simclr ; UEL ; he2019moco ; Insdis ; misra2020pirl to generate the corresponding pairs. This type of data augmentation has been shown critical for the success of CL, with different augmentation approaches having a different impact on SSL performance simclr . For video datasets, positive pairs are usually derived from temporal coherence constraints SermanetLHL17 ; Han19dpc . For multi-view data, augmentations can be more elaborate. For example, tian2019contrastive considers augmentations of color channels, depth or surface normal and shows that the performance of CL improves as the augmentations increase in diversity. Multi-modal CL approaches tend to rely on audio and video from a common video clip to design positive pairs morgado2020avid . In general, CL benefits from definitions of positive pairs that pose a greater challenge to the learning of the invariant instance representation.
Unlike the plethora of positive pair selection proposals, the design of negative pairs has received less emphasis in the CL literature. Since CL resembles approaches such as noise contrastive estimation (NCE) gutmann10a and N-pair npair losses, it can leverage hard negative pair mining schemes proposed for these approaches advconstest ; npair . However, because they treat each instance independently, previous SSL works do not consider CL algorithms that select difficult negative pairs within a batch. This is unlike CL methods based on metric learning KumarHC0D17 ; npair ; FaceNet ; Suh_2019_CVPR , which seek to construct batches with challenging negative samples.
In this work, we seek a general algorithm for the generation of diverse positive and challenging negative pairs. This is framed as the search for instance augmentation sets that induce the largest optimization cost for CL. A natural approach to synthesize these sets is to leverage adversarial examples adv_survey1 ; adv_survey2 ; fgsm ; Attackrcnn ; XieWZZXY17 ; XieWZZXY17 , which are crafted to attack the network and can thus be seen as the most challenging examples that it can process. We note that the goal is not to enhance robustness to adversarial attacks, but to produce a better representation for SSL. This is in line with recent studies in the adversarial literature, showing that adversarial examples can be used to improve supervised Xie2019AdversarialEI ; advprop_at_scale and semi-supervised vat learning. We explore whether the same benefits can ensue for SSL. One difficulty is, however, that no attention has been previously devoted to the design of adversarial attacks for SSL, where no class labels are available. In fact, for SSL embeddings trained by CL, the classical definition of adversarial attack does not even apply, since CL operates on pairs of examples. We show, however, that it is possible to leverage the interpretation of CL as instance classification to produce a sensible generalization of classification attacks to the CL problem. The new attacks are then combined with recent techniques from the adversarial literature Xie2019AdversarialEI ; advprop_at_scale , which treat adversarial training as multi-domain training, to produce more invariant representations for SSL.
Overall, the paper makes three contributions. First, we show that adversarial data augmentation can be used to improve the performance of SSL learning. Second, we propose a novel procedure for training Contrastive Learning with Adversarial Examples (CLAE) for SSL models. Unlike the attacks classically developed in the supervised learning literature, the new attacks produce pairs of examples that account for both the positive and negative pairs in a batch to maximize contrastive loss. To the best of our knowledge neither the use of attacks to improve SSL nor the design of adversarial examples with this property have been previously discussed in the literature. Finally, extensive experiments demonstrate that (1) adversarial examples can indeed be leveraged to improve CL, and (2) CLAE boosts the performance of several CL baselines across different datasets.
Related work
Since this work focuses on image classification tasks, our survey of previous work concentrates on contrastive learning (CL) and adversarial examples for image classification.
Contrastive learning has been widely used in the metric learning literature Chopra05 ; weinberger09a ; FaceNet and, more recently, for self-supervised learning (SSL) cpc ; Insdis ; UEL ; tian2019contrastive ; he2019moco ; simclr ; misra2020pirl ; TCN ; TCL , where it is used to learn an encoder in the pretext training stage. Under the SSL setting, where no labels are available, CL algorithms aim to learn an invariant representation of each image in the training set. This is implemented by minimizing a contrastive loss evaluated on pairs of feature vectors extracted from data augmentations of the image. While most CL based SSL approaches share this core idea, multiple augmentation strategies have been proposed Insdis ; UEL ; tian2019contrastive ; he2019moco ; simclr ; misra2020pirl . Typically, augmentations are obtained by data transformation (i.e. rotation, cropping, random grey scale and color jittering) UEL ; simclr , but there have also been proposals to use different color channels, depth, or surface normals as the augmentations of an image tian2019contrastive . Another approach is to use an augmentation dictionary composed of the embedding vectors from the previous epoch Insdis or obtained by forwarding an image through a momentum updated encoder he2019moco . This diversity of approaches to the synthesis of augmentations reflects the critical importance of using semantically similar example pairs in CL contrast_theory . This has also been studied empirically in simclr , showing that stronger data augmentations improve CL performance.
Despite this wealth of augmentation proposals for SSL, most CL algorithms fail to mine hard negative pairs or relate the image instances within a batch. While Wang15 ; MisraZH16 ; simclr ; Tschannen2020OnMI have mentioned the importance of selecting negative pairs, they do not propose a systematic algorithm to do this. Since CL is inspired by the noise contrastive estimation (NCE) gutmann10a and N-pair npair loss methods from metric learning, it inherits the well known difficulties of hard negative mining in this literature WuMSK17 ; npair ; FaceNet ; Suh_2019_CVPR ; KumarHC0D17 ; 7410379 ; BMVC2015_41 ; SongXJS15 . For metric learning methods FaceNet ; SongXJS15 ; KumarHC0D17 ; npair , the number of possible positive and negative pairs increases dramatically (for example, cubically when the triplet loss is used FaceNet ) as the dataset increases. A solution used by NCE is to draw negative samples from a noise distribution that treats all negative samples equallyadvconstest ; Bose2018CompositionalHN .
Unlike all these prior efforts, this work proposes to use “adversarial augmentations" as challenging training pairs that maximize the contrastive loss. However, unlike hard negative mining in metric learning, no class labels are provided in SSL. Hence, the consideration of how all images in the batch relate to each other is necessary for generating hard negative pairs.
2 Adversarial examples
Adversarial examples are created from clean examples to produce adversarial attacks that induce a network in error adv_survey1 ; adv_survey2 ; fgsm . They have been used in many supervised learning scenarios, including image classification fgsm ; KurakinGB16 ; PapernotMJFCS15 ; Moosavi_Dezfooli15 ; CarliniW16a ; Moosavi_Dezfooli16 ; onepixelattack , object detection Attackrcnn ; XieWZZXY17 ; Kevin18 ; Yue18 and segmentation XieWZZXY17 ; arnab_cvpr_2018 . Typically, to defend against such attacks and increase network robustness, the network is trained with both clean and adversarial examples, a process referred as adversarial training Ali19 ; tramer2018ensemble ; KurakinGB16a ; Lee2020AdversarialVM ; adv_survey1 ; advprop_at_scale . It is also possible to leverage SSL to increase robustness against unseen attacks Naseer_2020_CVPR . While adversarial training is usually effective as a defense mechanism, there is frequently a decrease in the accuracy of the classification of clean examples adv_trade_off ; Florian19 ; Tsipras2018RobustnessMB ; Xie2019AdversarialEI ; Lee2020AdversarialVM .
This effect has been attributed to overfitting to the adversarial examples Lee2020AdversarialVM but remains somewhat of a paradox, since the increased diversity of adversarial examples could, in principle, improve standard training adv_bug , e.g. by enabling models trained on adversarial examples to generalize better to unseen data volpi2018generalizing ; liu19b . In summary, while adversarial examples could assist learning, it remains unclear how to do this. Recently, Xie2019AdversarialEI ; advprop_at_scale have made progress along this direction, by introducing a procedure, denoted AdvProp, that treats clean and adversarial examples as samples from different domains, and uses a different set of batch normalization (BN) layers for each domain. This aims to align the statistics of the embedding of clean and adversarial samples, such that both can contribute to the network learning, and has been previously shown successful for multi-domain classification problems Bilen17 ; Rebuffi17 .
The proposed framework CLAE is inspired by recent advances in the adversarial example literature, yet parallel to them. Unlike these methods, we aim to leverage the strength of adversarial example for SSL, where no labels are available, and the focus on pairs rather than single examples requires an altogether different definition of adversaries. Our aims is to use adversarial training to compensate the limitation of current CL algorithms, by both generating challenging positive pairs and mining effective hard negative pairs for the optimization of the contrastive loss. Note that the goal is to produce better embeddings for CL, rather than robustifying CL embeddings against attacks.
Leveraging adversarial examples for improved contrastive learning
In this section, we introduce the approach proposed to create adversarial examples for CL, and a novel training scheme CLAE that leverages these examples for improved contrastive learning (CL).
A classifier maps example into label , where is a number of classes. A deep classifer is implemented by the combination of an embedding of parameters and a softmax regression layer that predicts the posterior class probabilities
where is the vector of classification parameters of class . Given a training set , the parameters and are learned by minimizing the risk defined by the cross-entropy loss
Given the learned classifier, the untargeted adversarial example of a clean example is
where is an adversarial perturbation of norm smaller than . The optimal perturbation for is usually found by maximizing the cross-entropy loss, i.e.
although different algorithms adv_survey1 ; adv_survey2 use different strategies to solve this optimization.
2 Contrastive learning
In SSL, the dataset is unlabeled, i.e. and each example is mapped into an example pair . In this work, we consider applications where is an image and the pair is generated by data augmentation. This consists of applying a transformation () in some set of transformations (e.g. spatial transformations, color transformations, etc.) to , to produce the augmentation (). CL seeks to learn an invariant representation of image by minimizing the risk defined by the loss
where is an embedding parameterized by , is the temperature, is the batch size and are augmentations of under transformations randomly sampled from . Previous works on CL have considered many possibilities for the set of transformations . While simclr has shown that the choice of has a critical role on SSL performance, most prior works do not give much consideration to the individual choice of and , which are simply uniformly sampled over . In this work, we seek to go beyond this and select optimal transformations for each image . More precisely, we seek augmentations that maximize the risk defined by the loss of (5), i.e.
This is, in general, an ill-defined problem since, for each example , what matters is the difference between the two transformations, not their absolute values.
3 Adversarial augmentation
This ambiguity can be eliminated by fixing one of the transformations of (5), i.e. solving instead
for a given set of sampled randomly from . However, it is usually difficult to search over the set efficiently. Instead, we proposed to replace by an adversarial perturbation of , i.e. find
where is a set of adversarial perturbations of , defined by
In summary, given a set of transformations and a set of augmentations , the goal is to learn the adversaries , by finding the perturbations that lead to the most diverse positive pairs (). The rationale is that the use of these pairs in (5) increases the challenge of unsupervised learning, encouraging the learning algorithm to produce a more invariant representation.
To optimize (8), we start by noting that the contrastive loss of (5) can be written as the cross-entropy loss of (2)
by defining the classifier parameters as . As above, we can replace each by the optimal adversarial perturbation , to obtain an optimal set of perturbed parameters by solving
Note that this requires the determination of the optimal adversarial perturbation for the augmentation of each example in the batch. Using the definition of adversarial set of (9) in (11), results in the optimization
This is an optimization similar to (4), but with a significant difference. While in (4) is a perturbation of the classifier input, in (12) it is a perturbation of the classifier weights. This implies that appears in the denominator of (10) for all in the batch and forces the optimization to account for all images simultaneously. In result, the embedding is more strongly encouraged to bring together the positive pairs and separate all perturbations of different examples, i.e. the optimization seeks both challenging positive and negative pairs, performing hard negative mining as well.
4 FSGM attacks for unsupervised learning
The optimization of (12) can be performed for any set of transformations . Hence, the procedure can be applied to most CL methods in the literature. The optimization can also be implemented with most adversarial techniques in the literature. In this work, we rely on untargeted attacks with the popular fast gradient sign method (FGSM) fgsm . For supervised learning, an untargeted FGSM attack consists of
where is the cross-entropy loss of (2). Similarly, to obtain of (12), the first order derivative of (12) is computed at to obtain with
Note that, due to this, the optimal set of adversarial augmentations takes into consideration the relationship between all instances within the sampled batch.
5 Adversarial training with contrastive loss
To perform adversarial training, we adopt the training scheme of AdvProp Xie2019AdversarialEI , which uses two separate batch normalization (BN) layers for clean and adversarial examples. Unlike Xie2019AdversarialEI ; advprop_at_scale , we set the momentum for the two BN layers differently. The momentum of the BN layer associated with clean examples is fixed to the value used by the original CL algorithm, while a larger momentum is used empirically for the BN layer associated with adversarial examples. The overall loss function is obtained by combining (10) and (15),
where balances between contrastive loss parametrized with and parametrized with . The proposed procedure for CLAE is summarized in Algorithm 1 and visualized in Fig.1. Given augmentations , adversarial examples are created by backpropagating the gradients of the contrastive loss to the network input. Contrastive loss terms are then computed for standard augmentation pairs and pairs composed by the augmentations and their adversarial examples, and combined with (16) to train the network.
6 Generalization to other contrastive learning approaches
While Algorithm 1 is based on the plain contrastive loss of (5), CLAE can be generalized to many other CL approaches in the literature. For example, UEL UEL uses an extra objective function to discourage different images from being recognized as the same instance. SimCLR simclr uses a projection head to avoid loss of information during pretext training and discards the projection head when optimizing the downstream task. To generalize Algorithm 1 to these approaches, it suffices to replace the plain contrastive loss with the loss functions on which they are based.
Experiments
In this section, we discuss an experimental evaluation of adversarial contrastive learning.
Experiments are performed on CIFAR10 cifar , CIFAR100 cifar or tinyImagenet tinyImagenet , using three different contrastive loss baselines: the loss of (5) (denoted as “Plain"), UEL UEL and SimCLR simclr . Unless otherwise noted, a Resnet18 encoder is trained using Algorithm 1 with , standard Pytorch augmentation, and an adversarial batchnorm momentum of 0.01Or, equivalently 0.99 for Tensorflow implementation. Two evaluation protocols are used, both based on a downstream classification task using features extracted by the learned encoder. These are implemented with a nearest neighbor (kNN) classifier, and a logistic regression layer (LR). The encoder is trained with batch size 256 (128) and LR is trained for 1000 (200) epochs for CIFAR10 and CIRFAR100 (tinyImagenet). See supplementary for more details.
2 Influence of adversarial example
In this section, we study the effect of adversarial examples on both the SSL surrogate and downstream tasks. We compare the adversarial attacks of Algorithm 1 to random additive perturbations of the same magnitude. Fig. 3(a) shows the magnitude of the contrastive loss of the pretext task, on CIFAR10. For a given perturbation strength , adversarial perturbations (blue) elicit a larger loss than random ones (red). Fig. 3(b) shows that, when the adversarial augmentations are fed into the downstream classification model, they elicit a larger cross entropy loss than random noise. This shows that adversarial augmentation produces more challenging pairs than random perturbations. Examples of adversarial augmentation are visualized in Fig. 2.
Nevertheless, there are some significant differences to the typical behaviour of adversarial examples in the supervised setting. First, the effect of adversarial attacks is weaker for SSL. In the SSL setting, adversarial examples only degrade the downstream classification accuracy by 20%. This is much weaker than previously reported for the supervised setting, using the adversarial examples of (4). Second, as illustrated in Fig. 4, the statistics of the batch normalization layers differ less between clean and adversarial examples than reported for supervised learning in advprop_at_scale . Both of these observations are explained by the fact that, while the classifier parameters of (2) remain stable across training, those of (10) vary between batches. This creates uncertainty in the perturbation direction of (15), decreasing the differences between clean and adversarial attacks under the SSL setting.
3 Comparison to contrastive learning baselines
In this section, we investigate the gains of using CLAE framework for SSL, and the consistency of these gains across CL approaches. Table 1 shows the downstream classification accuracy for the two classifiers and three CL methods considered in these experiments. Note that adversarial training reduces to these methods when , in which case no adversarial examples are used. It can be seen that adversarial training improves the performance of all CL algorithms on all datasets, with a gain that is consistent across downstream classifiers. While best performance is achieved by using in (15), still beats the baseline in most cases.
4 Ablation study
An ablation study was conducted using the combination of SimCLR simclr and LR classifier. The study was performed on both CIFAR100 and tinyImagenet, with qualitatively identical results. We present CIFAR100 results here and tinyImagenet on the supplementary.
Batch size Fig. 5(a) shows that larger batch sizes improve the performance of both baseline and adversarial training. While this is consistent with the conclusions of previous studies on the importance of batch size simclr ; khosla2020supervised , adversarial training makes large batch sizes less critical. Besides outperforming the baseline for all batch sizes, it frequently achieves better performance with smaller sizes. For example, adversarial training with and batch size 64 (54) outperforms the baseline with batch size 256 (53.79). This suggests that adversarial training is both more robust and efficient.
Embedding dimension impact is evaluated in Fig. 5(b). This dimension does not seem to affect the performance of both baseline and adversarially trained model.
Architecture The impact of different architectures is evaluated in Fig. 5(c), showing that larger networks have better performance. However, the improvement is much more dramatic with adversarial training (red bar), where it can be as high as 5 (ResNet101 over ResNet18), than for the baseline (green bar), where it is at most . This is likely because larger networks have higher learning capacity and can benefit from the more challenging examples produced by adversarial augmentation. Nevertheless, the fact that the ResNet18 with adversarial training beats the ResNet101 baseline shows that there is always a benefit to more challenging training pairs, even for smaller models.
Hyperparameter weights the contributions of the two losses ( and ) of (16). Fig. 6(a) shows downstream classification accuracy when varies from 0 to 2, for . Accuracy starts to increase with and reaches its peak around , but performance is fairly stable for . There is no significant accuracy drop when the importance of adversarial examples is weighted by as much as twice that of clean examples (). This shows that the features derived from the adversarial examples benefit learning.
Attack strength is studied in Fig. 6(b), which shows that adversarial training consistently beats the baseline (). Again, the performance is quite stable with . This is possibly due to the affect of batch normalization, which aligns the statistics of embeddings from clean and adversarial examples.
While the default setting of SimCLR is to use 100 training epochs, there are usually gains in using longer pretext training, as shown in Fig. 6(c). While all methods benefit from this, adversarial training consistently beats the baseline. It is also observed that the larger perturbation of benefits more from longer pretext training than smaller perturbations. Finally, training with adversarial augmentation is more efficient. At 300 epochs, it achieves results close to those of the baseline with 400 epochs; with 400 epochs, it outperforms the baseline at 500 epochs. This is similar to previous observations in metric learning, where challenging training pairs are known to improve the both convergence speed and final embedding performance.
Transfer to other downstream datasets Transfer performance compares how encoders learned by different SSL approaches generalize to various downstream datasets. Following the linear evaluation protocol of simclr , we consider the 8 datasets cifar ; cars ; aircraft ; dtd ; pets ; caltech101 ; flower shown in Table4.4. Both the encoder of SimCLR simclr and CLAE are trained on ImageNet100 (an ImageNet subset sampled by tian2019contrastive ), using a ResNet18. On ImageNet100, CLAE achieved 62.40.02 classification accuracy, outperforming SimCLR (61.70.02). This indicates that it can scale to large datasets. On the remaining datasets of Table4.4, it outperformed SimCLR on 7 out of the 8 datasets. This suggests that it generalizes better across downstream datasets. Since ImageNet100 does not contain any classes related to cars and airplanes, the performance on these 2 datasets is worse than on the others. In any case, CLAE beat simclr on several fine-grained datasets, such as Cars cars , Aircraft aircraft and Flowers flower .
Attack methods While FGSM fgsm is used in the above experiments, CLAE can be integrated with multiple attack approaches. To demonstrate this property, various attack methods (R-FGSM tramer2018ensemble , F-FGSM Wong2020Fast , PGD madry2018towards ) were evaluated, with , on CIFAR100. As shown in Figure 7, R-FGSM and F-FGSM have performance comparable to FGSM, while PGD is slightly weaker. However, all these attack methods beat the baseline, indicating that the proposed framework learns a better representation.
Conclusion
In self-supervised learning (SSL), approaches based on contrastive learning (CL) do not necessarily optimize on hard negative pairs. In this work, we have proposed a new algorithm (CLAE) that generates more challenging positive and hard negative pairs, on-the-fly, by leveraging adversarial examples. Adversarial training with the proposed adversarial augmentations was demonstrated to improve performance of several CL baselines. We hope this work will inspire further research on the use of adversarial examples in SSL.
Broader Impact
This work advances the general use of deep learning technology, especially in the case that dataset annotations are difficult to obtain, and could have many applications. It advances several state of the art solutions on self-supervised learning (SSL), where no labels are provided. Moreover, while prior works in SSL suggest training with larger network, larger batch size and longer training epochs, the experiments in this works demonstrates that these factors are less critical by optimizing on effective training pairs. This can be beneficial in the scenario where time and gpu resource are limited. While this work mainly focuses on the study of image recognition, we hope this work can be extended to other application domains of SSL in the future.
Acknowledgement
This work was partially funded by NSF awards IIS-1637941, IIS-1924937, and NVIDIA GPU donations. We also acknowledge and thank the use of the Nautilus platform for some of the experiments discussed above.