Generalizing to Unseen Domains via Adversarial Data Augmentation
Riccardo Volpi, Hongseok Namkoong, Ozan Sener, John Duchi, Vittorio Murino, Silvio Savarese
Introduction
In many modern applications of machine learning, we wish to learn a system that can perform uniformly well across multiple populations. Due to high costs of data acquisition, however, it is often the case that datasets consist of a limited number of population sources. Standard models that perform well when evaluated on the validation dataset—usually collected from the same population as the training dataset—often perform poorly on populations different from that of the training data Daume2006 ; Blitzer2006 ; BenDavid2006 ; Saenko2010 ; NameTheDataset . In this paper, we are concerned with generalizing to populations different from the training distribution, in settings where we have no access to any data from the unknown target distributions. For example, consider a module for self-driving cars that needs to generalize well across weather conditions and city environments unexplored during training.
A number of authors have proposed domain adaptation methods (for example, see Ganin ; ADDA ; DeepCORAL ; morerio2018 ; DIFA ) in settings where a fully labeled source dataset and an unlabeled (or partially labeled) set of examples from fixed target distributions are available. Although such algorithms can successfully learn models that perform well on known target distributions, the assumption of a priori fixed target distributions can be restrictive in practical scenarios. For example, consider a semantic segmentation algorithm used by a robot: every task, robot, environment and camera configuration will result in a different target distribution, and these diverse scenarios can be identified only after the model is trained and deployed, making it difficult to collect samples from them.
In this work, we develop methods that can learn to better generalize to new unknown domains. We consider the restrictive setting where training data only comes from a single source domain. Inspired by recent developments in distributionally robust optimization and adversarial training Certifiable ; LeeRa17 ; Heinze-DemlMe17 , we consider the following worst-case problem around the (training) source distribution
The solution to worst-case problem (1) guarantees good performance against data distributions that are distance away from the source domain . To allow data distributions that have different support to that of the source , we use Wasserstein distances as our metric . Our distance will be defined on the semantic space By semantic space we mean learned representations since recent works perceptual1 ; perceptual2 suggest that distances in the space of learned representations of high capacity models typically correspond to semantic distances in visual space., so that target populations satisfying represent realistic covariate shifts that preserve the same semantic representation of the source (e.g., adding color to a greyscale image). In this regard, we expect the solution to the worst-case problem (1)—the model that we wish to learn—to have favorable performance across covariate shifts in the semantic space.
We propose an iterative procedure that aims to solve the problem (1) for a small value of at a time, and does stochastic gradient updates to the model with respect to these fictitious worst-case target distributions (Section 2). Each iteration of our method uses small values of , and we provide a number of theoretical interpretations of our method. First, we show that our iterative algorithm is an adaptive data augmentation method where we add adversarially perturbed samples—at the current model—to the dataset (Section 3). More precisely, our adversarially generated samples roughly correspond to Tikhonov regularized Newton-steps Levenberg44 ; Marquardt63 on the loss in the semantic space. Further, we show that for softmax losses, each iteration of our method can be thought of as a data-dependent regularization scheme where we regularize towards the parameter vector corresponding to the true label, instead of regularizing towards zero like classical regularizers such as ridge or lasso.
From a practical viewpoint, a key difficulty in applying the worst-case formulation (1) is that the magnitude of the covariate shift is a priori unknown. We propose to learn an ensemble of models that correspond to different distances . In other words, our iterative method generates a collection of datasets, each corresponding to a different inter-dataset distance level , and we learn a model for each of them. At test time, we use a heuristic method to choose an appropriate model from the ensemble.
We test our approaches on a simple digit recognition task, and a more realistic semantic segmentation task across different seasons and weather conditions. In both settings, we observe that our method allows to learn models that improve performance across a priori unknown target distributions that have varying distance from the original source domain.
The literature on adversarial training FastGradientMethod ; Certifiable ; LeeRa17 ; Heinze-DemlMe17 is closely related to our work, since the main goal is to devise training procedures that learn models robust to fluctuations in the input. Departing from imperceptible attacks considered in adversarial training, we aim to learn models that are resistant to larger perturbations, namely out-of-distribution samples. Sinha et al. Certifiable proposes a principled adversarial training procedure, where new images that maximize some risk are generated and the model parameters are optimized with respect to those adversarial images. Being devised for defense against imperceptible adversarial attacks, the new images are learned with a loss that penalizes differences between the original and the new ones. In this work, we rely on a minimax game similar to the one proposed by Sinha et al. Certifiable , but we impose the constraint in the semantic space, in order to allow our adversarial samples from a fictitious distribution to be different at the pixel level, while sharing the same semantics.
There is a substantial body of work on domain adaptation Daume2006 ; Blitzer2006 ; Saenko2010 ; Ganin ; ADDA ; DeepCORAL ; morerio2018 ; DIFA , which aims to better generalize to a priori fixed target domains whose labels are unknown at training time. This setup is different from ours in that these algorithms require access to samples from the target distribution during training. Domain generalization methods DG0 ; DG1 ; DG2 ; DG3 ; Mancini2018 that propose different ways to better generalize to unknown domains are also related to our work. These algorithms require the training samples to be drawn from different domains (while having access to the domain labels during training), not a single source, a limitation that our method does not have. In this sense, one could interpret our problem setting as unsupervised domain generalization. Tobin et al. DomainRandomization proposes domain randomization, which applies to simulated data and creates a variety of random renderings with the simulator, hoping that the real world will be interpreted as one of them. Our goal is the same, since we aim at obtaining data distributions more similar to the real world ones, but we accomplish it by actually learning new data points, and thus making our approach applicable to any data source and without the need of a simulator.
Hendrycks and Gimpel SoftmaxICLR2016 suggest that a good empirical way to detect whether a test sample is out-of-distribution for a given model is to evaluate the statistics of the softmax outputs. We adapt this idea in our setting, learning ensemble of models trained with our method and choosing at test time the model with the greatest maximum softmax value.
Method
The transportation cost takes value for data points with different labels, since we are only interested in perturbation to the marginal distribution of . We now define our notion of distance on the semantic space. For inputs coming from the original space , we consider the transportation cost defined with respect to the output of the last hidden layer
so that measures distance with respect to the feature mapping . For probability measures and both supported on , let denote their couplings, meaning measures with and . Then, we define our notion of distance by
Armed with this notion of distance on the semantic space, we now consider a variant of the worst-case problem (1) where we replace the distance with (3), our adaptive notion of distance defined on the semantic space
Computationally, the above supremum over probability distributions is intractable. Hence, we consider the following Lagrangian relaxation with penalty parameter
Taking the dual reformulation of the penalty relaxation (4), we can obtain an efficient solution procedure. The following result is a minor adaptation of (BlanchetMu16, , Theorem 1); to ease notation, let us define the robust surrogate loss
In order to solve the penalty problem (4), we can now perform stochastic gradient descent procedures on the robust surrogate loss . Under suitable conditions BoydVa04 , we have
Iterative Procedure
We propose an iterative training procedure where two phases are alternated: a maximization phase where new data points are learned by computing the inner maximization problem (5) and a minimization phase, where the model parameters are updated according to stochastic gradients of the loss evaluated on the adversarial examples generated from the maximization phase. The latter step is equivalent to stochastic gradient steps on the robust surrogate loss , which motivates its name. The main idea here is to iteratively learn "hard" data points from fictitious target distributions, while preserving the semantic features of the original data points.
Concretely, in the -th maximization phase, we compute adversarially perturbed samples at the current model
where are the original samples from the source distribution . The minimization phase then performs repeated stochastic gradient steps on the augmented dataset . The maximization phase (8) can be efficiently computed for smooth losses if is strongly convex (Certifiable, , Theorem 2); for example, this is provably true for any linear network. In practice, we use gradient ascent steps to solve for worst-case examples (8); see Algorithm 1 for the full description of our algorithm.
Ensembles for classification
The hyperparameter —which is inversely proportional to , the distance between the fictitious target distribution and the source—controls the ability to generalize outside the source domain. Since target domains are unknown, it is difficult to choose an appropriate level of a priori. We propose a heuristic ensemble approach where we train models . Each model is associated with a different value of , and thus to fictitious target distributions with varying distances from the source . To select the best model at test time—inspired by Hendrycks and Gimpel SoftmaxICLR2016 —given a sample , we select the model with the greatest softmax score
Theoretical Motivation
We now give an interpretation for the augmented data points in the maximization phase (8). Concretely, we fix , , , and consider an -maximizer
Then, we have the following bound (10) whose proof we defer to Appendix A.1.
2 Data-Dependent Regularization
In this section, we argue that under suitable conditions on the loss,
The expansion (11) shows that the robust surrogate (5) is roughly equivalent to data-dependent regularization where we minimize the distance between , our “average estimated linear classifier”, to , the linear classifier corresponding to the true label . Concretely, for any fixed , we have the following result where we use to ease notation. See Appendix A.3 for the proof.
Experiments
We evaluate our method for both classification and semantic segmentation settings, following the evaluation scenarios of domain adaptation techniques Ganin ; ADDA ; FCNInTheWild , though in our case the target domains are unknown at training time. We summarize our experimental setup including implementation details, evaluation metrics and datasets for each task.
We train on MNIST MNIST dataset and test on MNIST-M Ganin , SVHN SVHN , SYN Ganin and USPS USPS . We use digit samples for training and evaluate our models on the respective test sets of the different target domains, using accuracy as a metric. In order to work with comparable datasets, we resized all the images to , and treated images from MNIST and USPS as RGB. We use a ConvNet ConvNet with architecture conv-pool-conv-pool-fc-fc-softmax and set the hyperparameters , , and . In the minimization phase, we use Adam Adam with batch size equal to Models were implemented using Tensorflow, and training procedures were performed on NVIDIA GPUs. Code is available at https://github.com/ricvolpi/generalize-unseen-domains. We compare our method against the Empirical Risk Minimization (ERM) baseline and different regularization techniques (Dropout Dropout , ridge).
Semantic scene segmentation
We use the SYTHIASYNTHIA dataset for semantic segmentation. The dataset contains images from different locations (we use Highway, New York-like City and Old European Town), and different weather/time/date conditions (we use Dawn, Fog, Night, Spring and Winter. We train models on a source domain and test on other domains, using the standard mean Intersection Over Union (mIoU) metric to evaluate our performance VOC2008 . We arbitrarily chose images from the left front camera throughout our experiments. For each one, we sample random images (resized to pixels) from the training set. We use a Fully Convolutional Network (FCN) FCN , with a ResNet-50 ResNet body and set the hyperparameters , , and . For the minimization phase, we use Adam Adam with batch size equal to . We compare our method against the ERM baseline.
1 Results on Digit Classification
In this section, we present and discuss the results on the digit classification experiment. Firstly, we are interested in analyzing the role of the semantic constraint we impose. Figure 1a (top) shows performances associated with models trained with Algorithm 1 with and , with the constraint in the semantic space (as discussed in Section 2) and in the pixel space Certifiable (blue and yellow bars, respectively). Figure 1a (bottom) shows performances of models trained with our method using different values of the hyperparameter (with ) and with ERM (blue bars and red lines, respectively). These plots show (i) that moving the constraint on the semantic space carries benefits when models are tested on unseen domains and (ii) that models trained with Algorithm 1 outperform models train with ERM for any value of on out-of-sample domains (SVHN, MNIST-M and SYN). The latter result is a rather desired achievement, since this hyperparameter cannot be properly cross-validated. On USPS, our method causes accuracy to drop since MNIST and USPS are very similar datasets, thus the image domain that USPS belongs to is not explored by our algorithm during the training procedure, which optimizes for worst case performance.
Figure 1b (top) reports results related to models trained with our method (blue bars), varying the number of iterations and fixing , and results related to ERM (red bars) and Dropout Dropout (yellow bars). We observe that our method improves performances on SVHN, MNIST-M and SYN, outperforming both ERM and Dropout Dropout statistically significantly. In Figure 1b (middle), we compare models trained with ridge regularization (green bars) with models trained with Algorithm 1 (with and ) and ridge regularization (blue bars); these results show that our method can potentially benefit from other regularization approaches, as in this case we observed that the two effects sum up. We further report in Appendix B a comparison between our method and an unsupervised domain adaptation algorithm (ADDA ADDA ), and results associated with different values of the hyperparameters and .
Finally, we report the results obtained by learning an ensemble of models. Since the hyperparameter is nontrivial to set a priori, we use the softmax confidences (9) to choose which model to use at test time. We learn ensemble of models, each of which is trained by running Algorithm 1 with different values of the as , with i=\big{\{}0,1,2,3,4,5,6\big{\}}. Figure 1b (bottom) shows the comparison between our method with different numbers of iterations and ERM (blue and red bars, respectively). In order to separate the role of ensemble learning, we learn an ensemble of baseline models each corresponding to a different initialization. We fix the number of models in the ensemble to be the same for both the baseline (ERM) and our method. Comparing Figure 1b (bottom) with Figure 1b (top) and Figure 1a (bottom), our ensemble approach achieves higher accuracy in different testing scenarios. We observe that our out-of-sample performance improves as the number of iterations gets large. Also in the ensemble setting, for the USPS dataset we do not see any improvement, which we conjecture to be an artifact of the trade-off between good performance on domains far away from training, and those closer.
2 Results on Semantic Scene Segmentation
We report a comparison between models trained with ERM and models trained with our method (Algorithm 1 with ). We set in every experiment, but stress that this is an arbitrary value; we did not observe a strong correlation between the different values of and the general behavior of the models in this case. Its role was more meaningful in the ensemble setting where each model is associated with a different level of robustness, as discussed in Section 2. In this setting, we do not apply the ensemble approach, but only evaluate the performances of the single models. The main reason for this choice is the fact that the heuristics developed to choose the correct model at test time in effect cannot be applied in a straightforward fashion to a semantic segmentation problem.
Figure 2 reports numerical results obtained. Specifically, leftmost plots report results associated with models trained on sequences from the Highway split and tested on the New York-like City and the Old European Town splits (top-left and bottom-left, respectively); rightmost plots report results associated with models trained on sequences from the New York-like City split and tested on the Highway and the Old European Town splits (top-right and bottom-right, respectively). The training sequences (Dawn, Fog, Night, Spring and Winter) are indicated on the x-axis. Red and blue bars indicate average mIoUs achieved by models trained with ERM and by models trained with our method, respectively. These results were calculated by averaging over the mIoUs obtained with each model on the different conditions of the test set. As can be observed, models trained with our method mostly better generalize to unknown data distributions. In particular, our method always outperforms the baseline by a statistically significant margin when the training images are from Night scenarios. This is since the baseline models trained on images from Night are strongly biased towards dark scenery, while, as a consequence of training over worst-case distributions, our models can overcome this strong bias and better generalize across different unseen domains.
Conclusions and Future Work
We study a new adversarial data augmentation procedure that learns to better generalize across unseen data distributions, and define an ensemble method to exploit this technique in a classification framework. This is in contrast to domain adaptation algorithms, which require a sufficient number of samples from a known, a priori fixed target distribution. Our experimental results show that our iterative procedure provides broad generalization behavior on digit recognition and cross-season and cross-weather semantic segmentation tasks.
For future work, we hope to extend the ensemble methods by defining novel decision rules. The proposed heuristics (9) only apply to classification settings, and extending them to a broad realm of tasks including semantic segmentation is an important direction. Many theoretical questions still remain. For instance, quantifying the behavior of data-dependent regularization schemes presented in Section 3 would help us better understand adversarial training methods in general.
References
Appendix A Proofs
Similarly as , let be an -optimizer to the problem (12)
by Assumption 1, where and denotes the maximum and minimum eigenvalue respectively. Recalling the definition of given in Eq (12), we then have
where we used the definition of in the last inequality.
Next, we note that and are close by Taylor expansion.
Using this inequality in the bound (14), we arrive at
From definition (13) of , we have
Next, to bound in the bound (15), we show that and are at most -away. We defer the proof of the following lemma to Appendix A.2
Applying Lemma 3 to bound on the right hand side of inequality (15), and using the bound (16) for , we obtain
A.2 Proof of Lemma 3
We use the following key lemma which says that for functions that satisfy a growth condition, its minimum is stable to perturbations to the function.
A.3 Proof of Theorem 2
Proof of Claim From Taylor’s theorem, we have
Using this approximation in the definition of , we get
Similarly, we can compute the lower bound
Combining the two bounds, the claim follows. ∎
Appendix B Additional Experimental Results
Table 1 reports results associated with the digit experiment (Section 4.1, Figure 2). In particular, it reports numerical results (averaged over different runs) obtained with models trained with Algorithm 1 by varying the hyperparameters and . Training set is constituted by MNIST samples, models were tested on SVHN, MNIST-M, SYN and USPS (see Figure 1 (top)). The baselines (accuracies achieved by models trained with ERM) are:
Table 2 reports results associated with the semantic segmentation experiment (Section 4.2, Figure 3). To summarize, it reports results obtained by training models on Highway and testing them on New York-like City and Old European Town, and by training models on New York-like City and testing them on Highway and Old European Town (see Figure 1 (bottom) to observe the different weather/time/date conditions). The comparison is between models trained with ERM (ERM rows) and our method (Ours rows), e.g.Algorithm 1 with and .
Finally, Figure 4 reports a comparison between our method (blue) and the unsupervised domain adaptation algorithm ADDA ADDA (yellow), by varying the number of target images fed to the latter during training. Note that, since unsupervised domain adaptation algorithms make use of target data during training while our method does not, the comparison is not fair. However, we are interested in evaluating to which extent our method can compete with a well performing unsupervised domain adaptation algorithm ADDA . While on MNIST USPS split ADDA clearly outperforms our method, on MNIST MNIST-M the accuracies reached by our method are just slightly lower than the ones reached by ADDA, and on MNIST SYN our method outperforms it, even if the domain adaptation algorithm has access to a large number of samples from the target domain. Finally, note that MNIST SVHN results are not provided because ADDA would not converge on this split (in effect, these results are neither reported in the original work ADDA ). Instead, models trained on MNIST samples using our method better generalize to SVHN, as shown in Section 4.1.