Learning to Generate Novel Domains for Domain Generalization

Kaiyang Zhou, Yongxin Yang, Timothy Hospedales, Tao Xiang

Introduction

Humans effortlessly generalize prior knowledge to novel scenarios, a capability that machines still struggle to reproduce. Typically, machine-learning models perform poorly when deployed on test data with a different data distribution than the training data, which is known as the domain shift problem . One line of research towards alleviating the domain shift problem is unsupervised domain adaptation (UDA), which exploits unlabeled target domain data for model adaptation . Although UDA methods avoid costly data annotation processes from target domains, data collection and per-domain model updates are still required. Meanwhile, UDA’s assumption that target data can be collected in advance is not always met in practice . This motivates another line of research, namely domain generalization (DG) , which is the main focus in this paper.

DG methods aim to learn models capable of good direct generalization to unseen target domains without data collection or model updating . They usually, but not always , leverage multiple source domains to train a generalizable model. Most existing DG methods focus on aligning available source domains , which is mainly inspired by UDA methods that seek to minimize the divergence between source data and unlabeled target data . As proved in , minimizing the domain divergence can lead to a smaller target error in the UDA setting. However, since DG methods focus on aligning source domains and do not have access to the target data, this theoretical proof does not apply to the DG setting. Recently, meta-learning has been exploited for DG where the key idea is to simulate domain shift by splitting the training data into meta-train and meta-test sets with non-overlapping domains . During learning, models are optimized on the meta-train domains in a way that the error is reduced on the meta-test domains. Nevertheless, similar to the alignment-based methods, meta-learning optimizes for reducing the domain gap among source domains, and thus still has the risk of overfitting to seen domains.

In this paper, we address DG from a different perspective, i.e., the most straightforward way to improve model generalization is increasing the diversity of available source domains (see Fig. 1). To this end, we propose L2A-OT (Learning to Augment by Optimal Transport). The core idea is to learn a conditional generator network that maps source domain images to pseudo-novel domains, and then combine both source and pseudo-novel domain images for training the actual task model. To train the generator, we maximize the distance between source domains and the generated pseudo-novel domains, as measured by optimal transport (OT) . This leads to the generated images having a very different distribution from the source domains (Fig. 1). However, this objective alone does not guarantee that the semantic content of the generated images is preserved. Therefore, we further impose two losses on the generator, namely a cycle-consistency loss and a classification loss, for maintaining the structural and semantic consistency respectively.

Our contributions are as follows. (1) For the first time, DG is tackled from a perspective of pseudo-novel domain synthesis. (2) A novel image generator is formulated which differs from existing generators in the objective (synthesizing pseudo-novel domain images vs. natural photo images). More importantly it has a unique OT-based formulation of objective functions that allow the generator to explore novel domain space and generate diverse data with distributions different from any of the original source domains. We evaluate L2A-OT on three homogeneous DG benchmark datasetsFollowing , homogeneous DG shares the same label space between training and test data while heterogeneous DG has disjoint label space. including digit recognition , PACS and Office-Home and a heterogeneous DG task in the form of cross-domain person re-identification (re-ID) . The results show that L2A-OT surpasses the current state-of-the-art on all datasets.

Related Work

Many DG methods are based on the idea of domain alignment popularized from the UDA literature , with a goal to learn a domain-invariant representation by minimizing the domain discrepancy between sources . As mentioned earlier, aligning domain distributions is mainly motivated by the theory developed for UDA, which does not apply to DG due to the absence of target data. Therefore, the models learned with domain alignment risk overfitting to source domains and as a result generalize poorly to unseen domains. In recent years, meta-learning has seen increasing interest for DG where the objective is to expose a model to domain shift during training. This can be achieved by dividing source domains into meta-train and meta-test sets without overlapping, and training a model on the meta-train set such that the error on the meta-test set is reduced . Similar to domain alignment methods, meta-learning methods still risk overfitting since the training data remains unchanged. Moreover, these methods work at feature-level, which is difficult for diagnosis and lacks visual interpretation. We refer readers to for a comprehensive DG survey.

Most related to our work are data augmentation methods, especially those based on adversarial gradients . For instance, proposed CrossGrad to perturb input images with adversarial gradients generated by a domain classifier. Different from adversarial gradient-based methods which only produce imperceptible and simple pixel-wise effects (due to the nature of adversarial attack ), our approach learns a full CNN generator to map source images to unseen domains and optimizes it via OT-based distribution divergence to make the new domains as dissimilar as possible to source distributions.

Domain randomization.

Our approach shares a similar high-level intuition with domain randomization (DR) , which was originally introduced in the context of robotic learning to improve generalization from simulation to real world. DR aims to diversify the training domains by changing the color and texture of objects, background scenes, lighting conditions, etc. via a computer simulator . Recently, DR has been successfully used in some computer vision applications, such as semantic segmentation and vehicle detection for autonomous driving . However, our approach is significantly different from the DR-based methods because we learn a CNN generator network from real images rather than using programmatic simulators. Thus our method is more scalable to a wider range of image recognition tasks.

Image-to-image translation.

Our work is also related to multi-domain image-to-image translation methods such as CycleGAN and StarGAN , which use GAN losses to generate realistic images and cycle-consistency losses to achieve translation without using paired training images. Our method is fundamentally different from CycleGAN/StarGAN in that our generator model is learned to map source images to unseen domains rather than performing mapping between source domains as did in CycleGAN/StarGAN. We show by experiments that simply doing source-to-source mapping for data augmentation offers little help to DG (see Fig. 5a).

Methodology

We are provided with KsK_{s} source domains with indices Ds={1,2,...,Ks}D_{s}=\{1,2,...,K_{s}\}. The goal is to learn a model which can generalize well on an unseen target domain. Without having access to the target data, we propose to improve the model’s generalization by synthesizing novel data domains Dn={1,2,...,Kn}D_{n}=\{1,2,...,K_{n}\} to augment the original source domains.

Conditional generator.

Objective functions.

In addition to maximizing the difference between source and novel distributions, we also maximize the difference between the generated novel distributions, i.e.

2 Maintaining Semantic Consistency

The model so far is optimizing a powerful CNN generator GG for the novelty of the generated distribution (Eq. (2) & (3)). This produces diverse images, but may not preserve their semantic content.

First, to guarantee structural consistency, we apply a cycle-consistency constraint to the generator,

Cross-entropy loss.

3 Training

The full objective for GG is the weighted combination of Eq. (2), (3), (4), & (5),

Task model training.

4 Design of Conditional Generator Network

Our generator model has a conv-deconv structure which is shown in Fig. 3. Specifically, the generator model consists of two down-sampling convolution layers with stride 2, two residual blocks and two transposed convolution layers with stride 2 for up-sampling. Following StarGAN , the domain indicator is encoded as a one-hot vector with length Ks+KnK_{s}+K_{n} (see Fig. 3). During the forward pass, the one-hot vector is first spatially expanded and then concatenated with the image to form the input to GG.

Though the design of GG is similar to the StarGAN model, their learning objectives are totally different: We aim to generate images that are different from the existing source domain distributions while the StarGAN model is trained to generate images from the existing source domains. In the experiment part we justify that adding novel-domain data is much more effective than adding seen-domain data for DG (see Fig. 5a). Compared with the gradient-based perturbation method in , our generator is allowed to model more sophisticated domain shift such as image style changes due to its learnable nature.

5 Design of Distribution Divergence Measure

Two common families for estimating the divergence between probability distributions are f-divergence (e.g., KL divergence) and integral probability metrics (e.g., Wasserstein distance). In contrast to most work that minimizes the divergence, we need to maximize it, as shown in Eq. (2) & (3). This strongly suggests to avoid f-divergence because of the near-zero denominators (they tend to generate large but numerically unstable divergence values). Therefore, we choose the second type, specifically the Wasserstein distance, which has been widely used in recent generative modeling methods .

The Wasserstein distance, also known as optimal transport (OT) distance, is defined as

where the soft-matching matrix MM represents the coupling distribution π\pi in Eq. (8) and can be efficiently computed using the Sinkhorn algorithm ; CC is the pairwise distance matrix computed over two sets of samples.

Following , we define the cost function as the cosine distance between instances,

where ϕ\phi is constructed by a CNN (also called critic in ), which maps images into a latent space. In practice, ϕ\phi is a fixed CNN that was trained with domain classification loss.

Experiments

(1) We use four different digit datasets including MNIST , MNIST-M , SVHN and SYN , which differ drastically in font style, stroke color and background. We call this new dataset Digits-DG hereafter. See Fig. 4a for example images. (2) PACS is composed of four domains, which are Photo, Art Painting, Cartoon and Sketch, with 9,991 images of 7 classes in total. See Fig. 4b for example images. (3) Office-Home contains around 15,500 images of 65 classes for object recognition in office and home environments. It has four domains, which are Artistic, Clipart, Product and Real World. See Fig. 4c for example images.

Evaluation protocol.

For fair comparison with prior work, we follow the leave-one-domain-out protocol in . Specifically, one domain is chosen as the test domain while the remaining domains are used as source domains for model training. The top-1 classification accuracy is used as performance measure. All results are averaged over three runs with different random seeds.

Baselines.

We compare L2A-OT with the recent state-of-the-art DG methods that report results on the same dataset or have code publicly available for reproduction. These include (1) CrossGrad , the most related work that perturbs input using adversarial gradients from a domain classifier; (2) CCSA , which learns a domain-invariant representation using a contrastive semantic alignment loss; (3) MMD-AAE , which imposes a MMD loss on the hidden layers of an autoencoder. (4) JiGen , which has an auxiliary self-supervision loss to solve the Jigsaw puzzle task ; (5) Epi-FCR , which designs an episodic training strategy; (6) A vanilla model trained by aggregating all source domains, which serves as a strong baseline.

Implementation details.

Results on Digits-DG.

Table 1 shows that L2A-OT achieves the best performance on all domains and consistently outperforms the vanilla baseline by a large margin. Compared with CrossGrad, L2A-OT performs clearly better on MNIST-M, SVHN and SYN, with clear improvements of 2.8%, 3.3% and 3%, respectively. It is worth noting that these three domains are very challenging with large domain variations compared with their source domains (see Fig. 4a). The huge advantage over CrossGrad can be attributed to L2A-OT’s unique generation of unseen-domain data using a fully learnable CNN generator, and using optimal transport to explicitly encourage domain divergence. Compared with the domain alignment methods, L2A-OT surpasses MMD-AAE and CCSA by more than 3.5% on average. The is because L2A-OT enriches the domain diversity of training data, thus reducing overfitting in source domains. L2A-OT clearly beats JiGen because the Jigsaw puzzle transformation does not work well on digit images with sparse pixels .

Results on PACS.

The results are shown in Table 2. Overall, L2A-OT achieves the best performance on all test domains. L2A-OT clearly beats the latest DG methods, JiGen and Epi-FCR. This is because our classifier benefits from the generated unseen-domain data while JiGen and Epi-FCR, like the domain alignment methods, are prone to overfitting to the source domains. L2A-OT beats CrossGrad on all domains, mostly with a large margin. This again justifies our design of learnable CNN generator over adversarial gradient.

Results on Office-Home.

The results are reported in Table 3. Again, L2A-OT achieves the best overall performance, and other conclusions drawn previously also hold. Notably, the simple vanilla model obtains strong results on this benchmark, which are even better than most existing DG methods. This is because the dataset is relatively large, and the domain shift is less severe compared with the style changes on PACS and the font variations on Digits-DG.

2 Evaluation on Heterogeneous DG

In this section, we evaluate L2A-OT on a more challenging DG task with disjoint label space between training and test data, namely cross-domain person re-identification (re-ID).

We use Market1501 and DukeMTMC-reID (Duke) . Market1501 has 32,668 images of 1,501 identities captured by 6 cameras (domains). Duke has 36,411 images of 1,812 identities captured by 8 cameras.

Evaluation protocol.

We follow the recent unsupervised domain adaptation (UDA) methods in the person re-ID literature and experiment with Market1501→\toDuke and Duke→\toMarket1501. Different from the UDA setting, we directly test the source-trained model on the target dataset without adaptation. Note that the cross-domain re-ID evaluation involves training a person classifier on source dataset identities. This is then transferred and used to recognize a disjoint set of people in the target domain of unseen camera views via nearest neighbor. Since the label space is disjoint, this is a heterogeneous DG problem. For performance measure, we adopt CMC ranks and mAP .

Implementation details.

Results.

In Table 4, we compare L2A-OT with the vanilla model and CrossGrad, as well as state-of-the-art UDA methods for re-ID. As a result, CrossGrad barely improves the vanilla model while L2A-OT achieves clear improvements on both settings. Notably, L2A-OT is highly competitive with the UDA methods, though the latter make the significantly stronger assumption of having access to the target domain data (thus gaining an unfair advantage). In contrast, L2A-OT generates images of unseen styles (domains) for data augmentation, and such more diverse data leads to learning a better generalizable re-ID model.

3 Ablation Study

To verify that our improvement is brought by the increase in training data distributions by the generated novel domains (i.e. Eq. (2) & (3)), we compare L2A-OT with StarGAN , which generates data from the existing source domains by performing source-to-source mapping. The experiment is conducted on Digits-DG and the average performance over all test domains is used for comparison. Fig. 5a shows that StarGAN performs only similarly to the vanilla model (StarGAN’s 73.8% vs. vanilla’s 73.7%) while L2A-OT obtains a clear improvement of 4.3% over StarGAN. This confirms that increasing domains is far more important than increasing data (of seen domains) for DG. Note that this 4.3% gap is attributed to the combination of the OT-based domain novelty loss (Eq. (2)) and the diversity loss (Eq. (3)). Fig. 5a shows that the diversity loss contributes around 1% to the performance, and the rest improvement comes from the diversity loss.

Importance of semantic constraint.

The cycle-consistency and cross-entropy losses (Eq. (4) & (5)) are essential in the L2A-OT framework for maintaining the semantic content when performing domain translation. Fig. 5b shows that without the semantic constraint, the content is completely missing (we found that using these images reduced the result from 78.1% to 73.9%).

4 Further Analysis

Our approach can generate an arbitrary number of novel domains, although we have always doubled the number of domains (set Ks=KnK_{s}=K_{n}) so far. Fig. 6 investigates the significance on the choice of number of novel domains. In principle, synthesizing more domains provides opportunity for more diverse data, but also increases optimization difficulty and is dependent on the source domains. The result shows that the performance is not very sensitive to the choice of novel domain number, with Kn=KsK_{n}=K_{s} being a good rule of thumb.

Do more source domains lead to a better result?

In general, yes. The evidence is shown in Table 5 where the result of using three sources is generally better than using two as we might expect due to the additional diversity. The detailed results show that when using two sources, performance is sensitive to the choice of sources among the available three. This is expected since different sources will vary in transferrability to a given target. However, for both vanilla and L2A-OT the performance of using three sources is better than the performance of using two averaged across the 2-source choices.

Visualizing domain distributions.

We employ t-SNE to visualize the domain feature embeddings using the validation set of Digits-DG (see Fig. 7a). We have the following observations. (1) The generated distributions are clearly separated from the source domains and evenly fill the unseen domain space. (2) The generated distributions form independent clusters (due to our diversity term in Eq. (3)). (3) GG has successfully learned to flexibly transform one source domain to any of the discovered novel domains.

Visualizing novel-domain images.

Fig. 8 visualizes the output of GG. In general, we observe that the generated images from different novel domains manifest different properties and more importantly, are clearly different from the source images. For example, in Digits-DG (Fig. 8a), GG tends to generate images with different background patterns/textures and font colors. In PACS (Fig. 8b), GG focuses on contrast and color. Fig. 8 seems to suggest that the synthesized domains are not drastically different from each other. However, a seemingly limited diversity in the image space to human eyes can be significant to a CNN classifier: both Fig. 1 and Fig. 7a show clearly that the synthesized data points have very different distributions from both the original ones and each other in a feature embedding space, making them useful for learning a domain-generalizable classifier.

L2A-OT vs. CrossGrad.

It is clear from Fig. 7b that the new domains generated by CrossGrad largely overlap with the original domains. This is because CrossGrad is based on adversarial attack methods , which are designed to make imperceptible changes. This is further verified in Fig. 9 where the images generated by CrossGrad have only subtle differences in contrast to the original images. On the contrary, L2A-OT can model much more complex domain variations that can materially benefit the classifier, thanks to the full CNN image generator and OT-based domain divergence losses.

Conclusion

We presented L2A-OT, a novel data augmentation-based DG method that boosts classifier’s robustness to domain shift by learning to synthesize images from diverse unseen domains through a conditional generator network. The generator is trained by maximizing the OT distance between source domains and pseudo-novel domains. Cycle-consistency and classification losses are imposed on the generator to further maintain the structural and semantic consistency during domain translation. Extensive experiments on four DG benchmark datasets covering a wide range of visual recognition tasks demonstrate the effectiveness and versatility of L2A-OT.

References