Domain Generalization via Model-Agnostic Learning of Semantic Features

Qi Dou, Daniel C. Castro, Konstantinos Kamnitsas, Ben Glocker

Introduction

Machine learning methods have achieved remarkable success, under the assumption that training and test data are sampled from the same distribution. In real-world applications, this assumption is often violated as conditions for data acquisition may change, and a trained system may fail to produce accurate predictions for unseen data with domain shift. To tackle this issue, domain adaptation algorithms normally learn to align source and target data in a domain-invariant discriminative feature space . These methods rely on access to a few labelled or unlabelled data samples from the target distribution during training.

An arguably harder problem is domain generalization, which aims to train a model using multi-domain source data, such that it can directly generalize to new domains without need of retraining. This setting is very different to domain adaptation as no information about the new domains is available, a scenario that is encountered in real-world applications. In the field of healthcare, for example, medical images acquired at different sites can differ significantly in their data distribution, due to varying scanners, imaging protocols or patient cohorts. At deployment, each new hospital can be regarded as a new domain but it is impractical to collect data each time to adapt a trained system. Learning a model which directly generalizes to new clinical sites would be of great practical value.

Domain generalization is an active research area with a number of approaches being proposed. As no a priori knowledge of the target distribution is available, the key question is how to guide the model learning to capture information which is discriminative for the specific task but insensitive to changes of domain-specific statistics. For computer vision applications, the aim is to capture general semantic features for object recognition. Previous work has demonstrated that this can be investigated through regularization of the feature space, e.g., by minimizing divergence between marginal distributions of data sources , or joint consideration of the class conditional distributions . Li et al. use adversarial feature alignment via maximum mean discrepancy. Leveraging distance metrics of feature vectors is another method . Model-agnostic meta-learning is a recent gradient-based method for fast adaptation of models to new conditions, e.g., a new task at few-shot learning. Meta-learning has been introduced to address domain generalization , by adopting an episodic training paradigm, i.e., splitting the available source domains into meta-train and meta-test at each iteration, to simulate domain shift. Promising performance has been demonstrated by deriving the loss from a task error , a classifier regularizer , or a predictive feature-critic module .

We introduce two complementary losses which explicitly regularize the semantic structure of the feature space via a model-agnostic episodic learning procedure. Our optimization objective encourages the model to learn semantically consistent features across training domains that may generalize better to unseen domains. Globally, we align a derived soft confusion matrix to preserve inter-class relationships. Locally, we use a metric-learning component to encourage domain-independent while class-specific cohesion and separation of sample features. The effectiveness of our approach is demonstrated with new state-of-the-art performance on two common object recognition benchmarks. Our method also shows consistent improvement on a medical image segmentation task. Code for our proposed method is available at: https://github.com/biomedia-mira/masf.

Related Work

Domain adaptation is based on the central theme of bounding the target error by the source error plus a discrepancy metric between the target and the source . This is practically performed by narrowing the domain shift between the target and source either in input space , feature space , or output space , generally using maximum mean discrepancy or adversarial learning . The success of methods operating on feature representations motivates us to optimize the semantic feature space for domain generalization in this paper.

Domain generalization aims to generalize models to unseen domains without knowledge about the target distribution during training. Different methods have been proposed for learning generalizable and transferable representations. A promising direction is to extract task-specific but domain-invariant features . Muandet et al. propose a domain-invariant component analysis method with a kernel-based optimization algorithm to minimize the dissimilarity across domains. Ghifary et al. learn multi-task auto-encoders to extract invariant features which are robust to domain variations. Li et al. consider the conditional distribution of label space over input space, and minimize discrepancy of a joint distribution. Motiian et al. use contrastive loss to guide samples from the same class being embedded nearby in latent space across data sources. Li et al. extend adversarial autoencoders by imposing maximum mean discrepancy measure to align multi-domain distributions. Instead of harmonizing the feature space, others use low-rank parameterized CNNs or decompose network parameters to domain-specific/-invariant components . Data augmentation strategies, such as gradient-based domain perturbation or adversarially perturbed samples demonstrate effectiveness for model generalization. A recent method with state-of-the-art performance is JiGen , which leverages self-supervised signals by solving jigsaw puzzles.

Meta-learning (a.k.a. learning to learn ) is a long standing topic exploring the training of a meta-learner that learns how to train particular models . Recently, gradient-based meta-learning methods have been successfully applied to few-shot learning, with a procedure purely leveraging gradient descent. The episodic training paradigm, originated from model-agnostic meta-learning (MAML) , has been introduced to address domain generalization . Epi-FCR alternates domain-specific feature extractors and classifiers across domains via episodic training, but without using inner gradient descent update. The method of MLDG closely follows the update rule of MAML, back-propagating the gradients from an ordinary task loss on meta-test data. A limitation is that using the task objective might be sub-optimal, as it is highly abstracted from the feature representations (only using class probabilities). Moreover, it may not well fit the scenario where target data are unavailable (as pointed out by Balaji et al. ). A recent method, MetaReg , learns a regularization function (e.g., weighted L1L_{1} loss) particularly for the network’s classification layer, excluding the feature extractor. Instead, Li et al. propose a feature-critic network which learns an auxiliary meta loss (producing a non-negative scalar) depending on output of the feature extractor. Both and lack notable guidance from semantics of feature space, which may contain crucial domain-independent ‘general knowledge’ for model generalization. Our method is orthogonal to previous work, proposing to enforce semantic features via global class alignment and local sample clustering, with losses explicitly derived in an episodic learning procedure.

Method

In the following, we denote input and label spaces by X\mathcal{X} and Y\mathcal{Y}, the domains D={D1,D2,…,DK}\mathcal{D}=\{D_{1},D_{2},\dots,D_{K}\} are different distributions on the joint space X×Y\mathcal{X}\times\mathcal{Y}. Since domain generalization involves a common predictive task, the label space is shared by all domains. In each domain, samples are drawn from a dataset Dk={(xn(k),yn(k))}n=1NkD_{k}=\{(\mathbf{x}_{n}^{(k)},y_{n}^{(k)})\}_{n=1}^{N_{k}} where NkN_{k} is the number of labeled data points in the kk-th domain. The domain generalization (DG) setting further assumes the existence of domain-invariant patterns in the inputs (e.g. semantic features), which can be extracted to learn a label predictor that performs well across seen and unseen domains. Unlike domain adaptation, DG assumes no access to observations from or explicit knowledge about the target distribution.

2 Global Class Alignment Objective

Relationships between class concepts exist in purely semantic space, independent of changes in the observation domain. In light of this, compared with individual hard label prediction, aligning class relationships can promote more transferable knowledge towards model generalization. This is also noted by Tzeng et al. in the context of domain adaptation, by aggregating the output probability distribution when fine-tuning the model on a few labelled target data. In contrast to their work, our goal is to structure the feature space itself to preserve learned class relationships on unseen data, by means of explicit regularization.

where Nk(c)N_{k}^{(c)} is the number of samples in domain Dk\mathcal{D}_{k} labelled as class cc. The obtained zˉc(k)\bar{\mathbf{z}}^{(k)}_{c} conveys how samples from a particular class are generally represented. It is then forwarded to the task network Tθ′T_{\theta^{\prime}}, for computing soft label distributions sc(k)\mathbf{s}^{(k)}_{c} with a ‘softened’ softmax at temperature τ>1\tau>1 :

3 Local Sample Clustering Objective

Contrastive loss is computed for pairs of samples, attracting samples of the same class and repelling samples of different classes . Instead of pushing clusters apart to infinity, the repulsion range is bounded by a distance margin ξ\xi.

Our contrastive loss for a pair of samples (n,m)(n,m) is defined as:

Triplet loss aims to make pairs of samples from the same class closer than pairs from different classes, by a certain margin ξ\xi . Given one ‘anchor’ sample aa, one ‘positive’ sample pp (with ya=ypy_{a}=y_{p}), and one ‘negative’ sample nn (with ya≠yny_{a}\neq y_{n}), we compute their triplet loss as follows:

Experiments

We evaluate and compare our method on three datasets: 1) the classic VLCS domain generalization benchmark for image classification, 2) the recently introduced PACS benchmark for object recognition with challenging domain shift, 3) a real-world medical imaging task of tissue segmentation in brain MRI. Results with an in-depth analysis and ablation study are presented in the following.

2 PACS Dataset

Results. Table 2 summarizes the results of object recognition on PACS dataset with a comparison to previous work (noting that not all compared methods reported results on both VLCS and PACS). MLDG and MetaReg employ episodic training with meta-learning, but from different angles in terms of the meta learner’s objective (Li et al. minimize task error, Balaji et al. learn a classifier regularizer). The promising results for indicate that exposing the training procedure to domain shift benefits model generalization to unseen domains. Our method further explicitly considers the semantic structure, regarding both global class alignment and local sample clustering, yielding improved accuracy. Across all domains, our method increases average accuracy by 3.51%3.51\% over the baseline. Note that current state-of-the-art JiGen improves 1.86%1.86\% over its own baseline. In addition, we observe an improvement of 6.20%6.20\% when the unseen domain is sketch, which has a distinct style and requires more general knowledge about semantic concepts.

Deeper architectures. In the interest of providing stronger baseline results, we perform additional preliminary experiments using more up-to-date deep residual architectures with ResNet-18 and ResNet-50. Table 4 shows strong and consistent improvements of MASF over the DeepAll baseline in all PACS splits for both network architectures. This suggests our proposed algorithm is also beneficial for domain generalization with deeper feature extractors.

3 Tissue Segmentation in Multi-site Brain MRI

Results. For easier comparison, we average the evaluated Dice scores achieved for the three foreground classes (GM/WM/CSF) and report it in Table 5. Although hard to notice visually from the gray-scale images, the domain shift from data distribution degrades segmentation significantly by up to 10%10\%. DeepAll is a strong baseline, yet our model-agnostic learning scheme provides consistent improvement over naively aggregating data from multiple sources, especially when generalizing to a new clinical site with relatively poorer imaging quality (i.e., Set-D). Figure 3 (c) is the Silhouette plot of the embeddings from MϕM_{\phi}, demonstrating that the samples within the same class cluster are tightly grouped, as well as clearly separated from those of other classes.

Test Train Set-A Set-B Set-C Set-D DeepAll MASF Set-A 90.62 88.91 88.81 85.03 89.09 89.82 Set-B 85.03 94.22 81.38 88.31 90.41 91.71 Set-C 93.14 92.80 95.40 88.68 94.30 94.50 Set-D 76.32 88.39 73.50 94.29 88.62 89.51

Conclusions

We have presented promising results for a new approach to domain generalization of predictive models by incorporating global and local constraints for learning semantic feature spaces. The better generalization capability is demonstrated by new state-of-the-art results on popular benchmarks and a dense classification task (i.e., semantic segmentation) for medical images. The proposed loss functions are generally orthogonal to other algorithms, and evaluating the benefit of their integration is an appealing future direction. Our learning procedure could also be interesting to explore in the context of generative models, which may greatly benefit from semantic guidance when learning low-dimensional data representations from multiple sources.

Acknowledgements

This project has received funding from the European Research Council (ERC) under the European Union’s Horizon 2020 research and innovation programme (grant No 757173, project MIRA, ERC-2017-STG) and is supported by an EPSRC Impact Acceleration Award (EP/R511547/1). DCC is also partly supported by CAPES, Ministry of Education, Brazil (BEX 1500/2015-05).

References