Progressive Domain Expansion Network for Single Domain Generalization

Lei Li, Ke Gao, Juan Cao, Ziyao Huang, Yepeng Weng, Xiaoyue Mi, Zhengze Yu, Xiaoya Li, Boyang xia

Introduction

In this paper, we define domains as various distributions of objects appearance caused by different external conditions(such as weather, background, illumination etc.) or intrinsic attributes(such as color, texture, pose etc.), as shown in Fig.1. The performance of a deep model usually drops when applied to unseen domains. For example, The accuracy of the CNN model(trained on MNIST) on MNIST test set is 99%, while that on SVHN test set is only 30%. Model generalization is important to machine learning.

Two solutions have been proposed to deal with the above issue, namely, domain adaptation and domain generalization . Domain adaptation aims to generalize to a known target domain whose labels are unknown. Distribution alignment(e.g., MMD) and style transfer(e.g., CycleGAN) are frequently used in these methods to learn domain-invariant features. However, it requires data from the target domain to train the model, which is difficult to achieve in many tasks due to lack of data.

Domain generalization, which not requires access to any data from the unseen target domain, can solve these problems. The idea of domain generalization is to learn a domain-agnostic model from one or multiple source domains. Particularity, in many fields we are usually faced with the challenge of giving a single source domain, which is defined as single domain generalization . Recently, studies have made progress on this task . All of these methods, which are essentially data augmentation, improve the robustness of the model to the unseen domain by extending the distribution of the source domain. Specifically, additional samples are generated by manually selecting the augmentation type or by learning the augmentation through neural networks.

Data augmentation has proved to be an important means for improving model generalization . However, such methods require the selection of an augmentation type and magnitude based on the target domain, which is difficult to achieve in other tasks. They cannot guarantee the safety and effectiveness of synthetic data or even reduce accuracy. .

In this paper, we propose the progressive domain expansion network (PDEN) to solve the single domain generalization problem. Task models and generators in PDEN mutually benefit from each other through joint learning. Safe and effective domains are generated by the generator under the precise guidance of the task model. The generated domains are progressively expanded to increase the coverage and improve the completeness. Contrastive learning is introduced to learn the cross-domain invariant representation with all the generated domains. It is noteworthy that we can flexibly replace the generator in PDEN to achieve different types of domain expansion.

We propose a novel framework called progressive domain expansion network (PDEN) for single domain generalization. The PDEN contains domain expansion subnetwork and domain invariant representation learning subnetwork, which mutually benefit from each other by joint learning.

For the domain expansion subnetwork, multiple domains are progressively generated to simulate various photometric and geometric transforms in unseen domains. A series of strategies are introduced to guarantee the safety and effectiveness of these domains.

For the domain invariant representation learning subnetwork, contrastive learning is introduced to learn the domain invariant representation in which each class is well clustered so that a better decision boundary can be learned to improve it’s generalization.

Extensive experiments on classification and segmentation have shown the superior performance of our method. The proposed method can achieve up to 15.28% improvement compared with other single-domain generalization methods.

Related Work

Domain Adaptation. In recent years, many domain adaptation methods have been proposed to solve the problem of domain drift between source and target domain, including feature-based adaptation, instance-based adaptation and model parameter based adaptation . The domain adaptation method in deep learning is mainly to align the distribution of source domain and target domain, including two kinds of methods: MMD based adaptation method and adversarial based method. DDC is first proposed to solve the domain adaptation problems in deep networks. DDC fixes the weights of the first 7 layers in AlexNet, and MMD is used on the 8th layer to reduce the distribution difference between the source domain and target domain. DAN increased the number of adaptive layers (three in front of the classifier head) and introduced MK-MMD instead of MMD. AdaBN proposed to measure the distribution of the source domain and target domain in BN layer. With the emergence of GAN, a lot of domain adaptation methods based on adversarial learning have been developed. DANN is the first research work to reduce the distribution difference between the source domain and target domain by adversarial learning. DSN assumes that each domain includes a domain-shared distribution and a domain-specific distribution. Based on this assumption, DSN learned the shared feature and the domain-specific feature respectively. DAAN measures the marginal distribution and conditional distribution with a learnable weight.

Domain Generalization. Domain generalization is more challenging than domain adaptation.Domain generalization aims to learn the model with data from the source domain and the model can be generalized to unseen domains.

Domain generalization can be categorized as such several research interests: Domain alignment and domain ensemble. Domain alignment methods assume that there is a distribution shared by different domains. These methods map the distribution from different domains to the shared one. CCSA propose the contrastive semantic alignment loss to minimize the distance between data with the same label but from different domains and maximize the distance between the data with different class labels. In MMD-AAE, the feature distribution of source domain and target domain are aligned by MMD, and then the feature representation is matched to the prior Laplace distribution by AAE. Model ensemble methods train models for each source domain in the training set, and then ensemble their outputs according to the confidence of each model.

Single-domain generalization assumes that the training set only contains samples from just one source domain. Recent, many studies have made progress on this task . These methods are generally applied to synthesize more samples in image space or feature space to expand the range of data distribution in the training set. BigAug observed that the differences in medical images (such as T2 MRI) are mainly different in 3 aspects: image quality, image appearance, and spatial configuration. They augment more variants for the 3 aspects by data augmentation. However, such methods require the selection of an augmentation type and magnitude based on the target domain, which is difficult to achieve in other tasks. GUD and MADA synthesize more data through adversarial learning to promote the model’s robustness. However, on the one hand, the augmentation type is relatively simple; on the other hand, too much adversarial examples used for training will damage the performance of the classifier.

Contrastive Learning. Contrastive learning is a kind of unsupervised pre-training method for image recognition, which is popular these years. The key idea of contrastive learning is to train a model by bringing the positive pairs closer and pushing apart negative pairs. SimCLR generates positive pairs by imposing strong augmentation on whole images. CPC utilizes augmentation on image patches and uses the patch-level views for loss optimization.

Method

The PDEN proposed in this paper is used for single domain generalization. Suppose the source domain is S={xi,yi}i=1NS\mathcal{S}=\{x_{i},y_{i}\}_{i=1}^{N_{S}}, the target domain is T={xi,yi}i=1NT\mathcal{T}=\{x_{i},y_{i}\}_{i=1}^{N_{T}}, where xi,yix_{i},y_{i} is the ii-th image and class label, NS,NTN_{S},N_{T} represent the number of samples in source domain and target domain respectively. The aim is to train the model with only S\mathcal{S} then it can be generalized to the unseen T\mathcal{T}.

The overall model architecture of PDEN is shown in Fig. 2 (b), including the task net MM and unseen domain generator GG. In this section, we will introduce the task model in the PDEN.

There are 3 parts in MM. (1) Feature extractor F:X→HF:\mathcal{X}\rightarrow\mathcal{H}, where X\mathcal{X} is the image space and H\mathcal{H} is the feature space. FF is a stack of convolution layers followed by the pooling layers and activation layers. The output of FF is a 1-d vector obtained by global pooling. (2) Classifier head C:H→YC:\mathcal{H}\rightarrow\mathcal{Y}, where Y\mathcal{Y} is the label space. Here we focus on the classification task, so the task head CC is optimized by cross-entropy loss. In our experiment, CC is a stack of fully connected layers followed by nonlinear activation layers, and the last activation layer in CC is softmax. (3) Projection head P:H→ZP:\mathcal{H}\rightarrow\mathcal{Z}, where Z\mathcal{Z} is the hidden space in which the contrastive loss will be calculated. PP contains only one full connection layer in our experiments. We normalize the output vector of PP to lie on a unit hypersphere, which enables using an inner product to measure similarity in the Z\mathcal{Z} space.

2 The Unseen Domain Generator G𝐺G

GG can convert the original image xx(original domain S\mathcal{S}) to a new image x^\hat{x}(unseen domain S^\hat{\mathcal{S}}) as follows:

where x^\hat{x} has the same semantic information as xx, but the domains of x^\hat{x} and xx is different.

GG can be a variety of structures depending on related downstream tasks, such as AutoEncoder , HRNet , spatial transform network(STN) or a combination of these networks.

Autoencoder as GG: In our experiment, we mainly use the Autoencoder with AdaIN as the generator, as the GkG_{k} shown in Fig 2. The generator GG contains the encoder GEG_{E}, the AdaIN and the decoder GDG_{D}. In AdaIN, there are two fully-connected layers Lfc1,Lfc2L_{fc1},L_{fc2}:

where n∼N(0,1)n\sim N(0,1). Fig.3(a) shows the unseen domains generated by Autoencoder.

STN as GG: The Autoencoder can be replaced by the STN as the generator. The STN is a geometry-aware module which can transform the spatial structure of the image. Fig.3(b) shows the unseen domains generated by STN.

PDEN is a framework in which generators can be replaced with different structures depending on the tasks. In our experiment, the autoencoder is applied.

3 Progressive Domain Expansion

In order to improve the completeness of the generated domains and expand its coverage, we progressively generate KK unseen domains {S^k=Gk(S)}k=1K\{\hat{\mathcal{S}}_{k}=G_{k}(\mathcal{S})\}_{k=1}^{K} with the learnable generator GG. The task model MM is trained with these unseen domains to learn the cross-domain invariant representation. We train the task model and generator alternately, as shown in Fig. 2.

Take the kthkth domain expansion as an example. First, the generator GG and task model MM are jointly trained to synthesize safe and effective unseen domains S^k\hat{\mathcal{S}}_{k} by minimize Equ.(9). Then, the task model MM will be retrained with the updated data set S∪{S^i}i=1k\mathcal{S}\cup\{\hat{\mathcal{S}}_{i}\}_{i=1}^{k} by minimize Equ.(3). The performance of MM will be improved, so MM can guide the generator Gk+1G_{k+1} to synthesize better unseen domains. The algorithm is shown in Alg.1.

4 Domain Alignment and Classification

In this section, we will introduce how to learn cross-domain invariant representation. Given a minibatch B={xi,yi}i=12N\mathcal{B}=\{x_{i},y_{i}\}_{i=1}^{2N}, xix_{i} is the source image, xi+=G(xi,n)x_{i}^{+}=G(x_{i},n) is the synthetic image originating from xix_{i}(xix_{i} and xi+x_{i}^{+} have the same semantic information, but come from different domains), yiy_{i} is the class label. MM is optimized by:

where yimy_{i}^{m} is the mthm_{th} dimension of yiy_{i}; y^i=C(F(xi))\hat{y}_{i}=C(F(x_{i})); zi=P(F(xi))z_{i}=P(F(x_{i})).

LceL_{ce} is the cross-entropy loss used for classification. LNCEL_{NCE} is the InfoNCE loss used for contrastive learning. In the minibatch B\mathcal{B}, ziz_{i} and zi+z_{i}^{+} have the same semantic information but come from different domains. By minimizing LNCEL_{NCE}, the distance between ziz_{i} and zi+z_{i}^{+} will be smaller. In other words, samples from different domains with the same semantic information will be closer in the Z\mathcal{Z} space. LNCEL_{NCE} will guide FF to learn domain-invariant representation.

5 Unseen Domain 𝒮^^𝒮\hat{\mathcal{S}} Generation

In this section, we will show how to generate kkth unseen domain Sk^\hat{\mathcal{S}_{k}} from S\mathcal{S} via the generator GkG_{k}(For convenience, we use G,S^G,\hat{\mathcal{S}} instead of Gk,S^kG_{k},\hat{\mathcal{S}}_{k}). S^\hat{\mathcal{S}} satisfy the constraints of safety and effectiveness. Safety means the generated samples contain the domain-invariant information. Effectiveness means the generated samples contain various unseen domain-specific information.

Safety. S^\hat{\mathcal{S}} is safe if all the x∈S^x\in\hat{\mathcal{S}} can be predicted correctly by task model MM. Formally, we optimize:

Cycle consistency loss is introduced to further ensure the safety of S^\hat{\mathcal{S}}. S^\hat{\mathcal{S}} is safe if it can be converted to S\mathcal{S} by an generator GcycG_{cyc}. GcycG_{cyc} has the same structure as GG, but no noise input. Formally, we optimize:

Effectiveness. Adversarial learning is introduced to generate effective unseen domains. The generator GG and task model MM are learned jointly. The task model MM which extracts the domain-share representation is always trained to minimize the InfoNCE loss. The generator GG is trained to maximize the InfoNCE loss. Through adversarial training, GG will generate unseen domains from which MM can’t extract domain shared representation, and MM will be better able to extract cross-domain invariant representations. The loss can be defined as:

We also use a loss function to encourage GG to generate more diverse samples.

where n1,n2∼N(0,1)n1,n2\sim N(0,1), and n1≠n2n1\not=n2. To sum up, the loss function of training generate GG is as follow:

The weight of LclsL_{cls} is always 1, wcyc,wadv,wdivw_{cyc},w_{adv},w_{div} are the weights of Lcyc,Ladv,LdivL_{cyc},L_{adv},L_{div}.

Experiment

Follow , we evaluated our approach on Digits, CIFAR10-C and SYNTHIA.

Digits Dataset: Digits dataset contains 5 datasets: MNIST, MNSIT-M, SVHN, USPS, SYNDIGIT. Each dataset is considered as a domain. We use MNIST as the source domain and the other four data sets as the target domains. The first 10,000 images in MNIST are used to train the model.

CIFAR10-C Dataset: We use the CIFAR10 as the source domain and the CIFAR10-C as the target domain. CIFAR10-C is a benchmark dataset to evaluate the robustness of classification models. CIFAR10-C dataset consists of test images with 19 corruption types, which are algorithmically generated. The corruptions come from 4 categories and each type of corruption has 5 levels of severity.

SYNTHIA Dataset: The SYNTHIA VIDEO SEQUENCES dataset is used for traffic scene segmentation. The dataset consists of 3 locations: Highway, New York ish and Old European Town. Each location consists of the same traffic situation but under different weather/illumination/season conditions(we use Dawn, Fog, Spring, Night and Winter in our experiment). Following the protocol in, we train our model on one domain and evaluate on the other domains. For each domain, we randomly sample 900 images from the left front camera and all the images are resized to 192×320192\times 320 pixels.

Evaluate: For Digit and CIFAR10 datasets, we compute the mean accuracy on each unseen domain. For SYNTHIA dataset, we use the standard mean Intersection over Union(mIoU) to evaluate the performance on each unseen domain.

2 Evaluation of Single Domain Generalization

We compare our method with the following state-of-the-art methods. (1) Empirical Risk Minimization(ERM) is the baseline method trained with only the cross-entropy loss. (2) CCSA aligns samples from different domains of the same category to get the robust feature space for domain generalization. (3) d-SNE minimizes the maximum distance between sample pairs of the same class and maximizes the minimum distance among sample pairs of different categories. (4) GUD proposes an adversarial data augmentation method to synthesize more hard samples which can improve the robustness of the classifier. (5) MADA minimizes the distance of semantic space and maximize the distance of pixel space to generate more effective samples. (6) JiGen proposes a multi-task learning method that combines the target recognition task and the Jigsaw classification task to improve the cross-domain generalization of the model. (7) AutoAugment(AA) proposes a method to automatically searches improved data augmentation policies for the specific data set. (8) Based on AA, RandAugment(RA) has a better data augment policies, which greatly reduces the policies space .

Comparison on Digits: We train the model with the first 10,000 images in the MNIST train set, validate the model on the MNIST test set, and evaluate the model on the MNIST-M, SVHN, USPS, and Syndigits datasets. We calculate the mean accuracy on each data set as the evaluation index. We first compared with the single-domain generalization methods, as shown in the top half of Table 1. To be fair, we did not use any manual data augmentation. We observed that our method performs much better than other methods on SVHN, MNIST-M and USPS. On USPS, the performance of our method is comparable to others, mainly because the USPS is more similar to MNIST. The d-SNE performs well on USPS, but bad on other data sets. We also compare with the data augmentation methods as shown in the bottom half of Table 1. The hyperparameters are consistent with those in the original paper. We found that our method performs better than these methods. What’s more, our approach is orthogonal to these data augmentation techniques.

Comparison on CIFAR10: We train all the models on the CIFAR10 train set, validate the models on the CIFAR10 test set, and evaluate the models on the CIFAR10-C. The experimental results across five levels of corruption severity are shown in Tab2. Our approach performs better than other single-domain generalization methods such as GUD and MADA. The severer the corruption, the more our approach surpasses MADA. Compared to approaches using manual data augmentation, our approach performs as well as they do at lower corruption levels and outperform them at higher corruption levels. We also show the experimental results across different types of corruptions with the 5th level severity in Tab 3. Our approach has higher average accuracy than other approaches. In some corruption types, the RandAugment approach performs better than us. However, it is important to note that there is no manual data augmentation in our approach, and our approach can be used together with RandAugment.

Comparison on SYNTHIA: Follow the protocol in , we conducted three experiments, using Highway-Dawn, Highway-Fog and Highway-Spring as the source domain respectively, and taking all the weather in New York ish and Old European Town as the unseen target domains. The scene segmentation results(mIoU) are show in Tab 4. Our approach improves the average mIoU compared to other approaches. When the source domain is highway-Dawn or Highway-fog, the improvement is greater.

3 Additional Analysis

Validation of KK: We study the effect of the hyper-parameters KK on the Digits dataset. We use the MNIST as the source domain, and take MNIST-M, SVHN, USPS and SYNDIGIT as the unseen target domains. The experimental result is shown in Fig.5(a). We report the classification accuracy on the target domains when K=1,2,...,20K=1,2,...,20. The accuracy is increased rapidly when KK is small, and gradually converges when KK is large. In experiments on Digits, we set K=20K=20. In the Digits experiment in MADA, their approach performed best at K=3K=3 and decreased as KK grew.This indicates that the domain generated by our approach is safer than MADA.

Validation of wadvw_{adv}: We study the effect of the hyper-parameters wadvw_{adv} on Digits dataset. The experimental results are shown in Fig.5(b). We report the classification accuracy on target domains when wadv=0.02,0.05,0.08,0.1,0.13,0.16,0.2w_{adv}=0.02,0.05,0.08,0.1,0.13,0.16,0.2. We found that the accuracy increases with the increase of wadvw_{adv} on the unseen target domains.

Validation of wcycw_{cyc}: We study the effect of the hyper-parameters wcycw_{cyc} on Digits dataset. The experimental result is shown in Fig.5(c). We report the classification accuracy on MNIST-M, USPS, SVHN and SYNDIGIT when wcyc=0,10,20,30,40,50w_{cyc}=0,10,20,30,40,50. On MNIST-M, SVHN and SYNDIGIT, the accuracy increases with the increases of wcycw_{cyc}. On USPS, the classification accuracy did not change significantly (fluctuated within a small range) with the increase of wcycw_{cyc}, mainly because of the high similarity between USPS and MNIST.

Validation of wdivw_{div}: We illustrate the effect of the hyper-parameters wdivw_{div} in Fig.5(d). For all the unseen domain in Digit dataset, the classification accuracy increases with the increases of wdviw_{dvi}.

Visualization of the feature space: Fig.4 illustrates the difference in 2-d feature spaces between PDEN and the baseline models. For PDEN, the sample distribution of target domain is consistent with that of source domain. For the baseline model, most of the target samples are mixed in the feature space so that it is difficult to classify them.

4 Evaluation of of Few-shot Domain Adaptation

We also compared our methods in the experimental setting of few-shot domain adaptation . In few-shot domain adaptation, data from the source domain S\mathcal{S} and a few samples from the target domain T\mathcal{T} are used to train the model.

We use MNIST as the source domain and SVHN as the target domain. We first train the model on mnist with the proposed PDEN, and then finetune the model with few samples from SVHN. The model will be evaluated on SVHN, as shown in Fig. 6. We found that finetuning with few samples from the target domain can significantly improve the performance of the model on the target domain. Compared with MADA, the proposed PDEN performs better in this case.

Conclusion

In this paper, we propose a single domain generalization learning framework to learn the domain-invariant feature, which can generalize the model to the unseen domains. We learn the generator to synthetic unseen domains, which share the same semantic information as the source domain. The domain-invariant representation can be learned by aligned the source and unseen domain distribution. We mine the hard unseen domains in which the domain-invariant representation cannot be extracted by the task model. The model will be more robust by adding these generated domains to the training set. The novel method PDEN proposed in this paper provides a promising direction to solve the single-domain generalization problem.

Acknowledgements

This work was supported by the National Key Research and Development Program of China (2018YFC0825202), and the National Natural Science Foundation of China (U1703261,61871004), and Beijing Municipal Natural Science Foundation Cooperation Beijing Education Committee: No. KZ 201810005002.

References