Learning to Learn with Variational Information Bottleneck for Domain Generalization
Yingjun Du, Jun Xu, Huan Xiong, Qiang Qiu, Xiantong Zhen, Cees G. M. Snoek, Ling Shao
Introduction
This paper strives for domain generalization in image classification . The general challenge is to exploit the data variations of seen image domains with the aim to generalize well to unseen image domains. For example, by generalizing a chair classifier trained on PASCAL VOC to LabelMe , or by generalizing an elephant classifier trained on photo’s to sketches . Domain generalization models typically suffer from two problems. First, since data from unseen domains is inaccessible during the learning stage, we do not know their statistical data distribution. This causes uncertainty in the predictions made on the unseen domains. Second, data from different domains usually follows distinct distributions with great discrepancy, resulting in domain shift from seen to unseen domains. Domain shift has been extensively researched in domain generalization, mostly by learning feature representations that are invariant across domains . Meta-learning that learns to generalize across tasks has been introduced to domain generalization by Li et al. showing its great effectiveness in learning to generalize across domains . To the best of our knowledge, none of these existing meta-learning methods deal with the prediction uncertainty on unseen domains.
In this paper, we address the two major domain generalization challenges jointly by one single probabilistic model under the meta-learning framework. We model parameters of classifiers shared across domains as probabilistic distributions that we infer from the data of the seen domains. The probabilistic modeling enables us to better handle the prediction uncertainty on previously unseen domains . To handle domain shift, we take inspiration from the information bottleneck (IB) theory which learns robust representations to enhance generalization. IB encodes the input into compressed intermediate representations that maximize target prediction. It offers a promising technique to learn domain-invariant representations, but to the best of our knowledge has not yet been explored for domain generalization under the meta-learning framework. We propose the principle of meta variational information bottleneck (MetaVIB) for the optimization of the model. We derive MetaVIB from the variational bounds of mutual information by leveraging the meta-learning setting, and incorporate it as a data-driven regularizer into the optimization objective. The parameters of all classifiers and the network are jointly optimized during the meta-training stage and applied to the unseen domain in the meta-test stage. By episodic training, MetaVIB enables the network to learn to gradually close the gaps between domains to achieve domain-invariant representations that alleviate domain shift, while simultaneously being able to obtain accurate predictions.
We conduct extensive experiments on three benchmarks for cross-domain visual recognition. The ablation studies demonstrate the benefits of MetaVIB in the probabilistic framework for domain generalization. The comparison with state-of-the-art methods, shows that our method consistently delivers the best performance on all tasks, surpassing previous methods based on both regular learning and meta-learning.
Related Work
In this section, we review related work on domain generalization, information bottleneck and meta-learning.
Domain generalization has been a longstanding challenge in computer vision and machine learning ,but recently regained increased research interest . Learning domain-variant feature representation has been one of the main topics of focus in domain generalization . The core idea is to learn a model that generates invariant representations for the source domains, without over-fitting, which generalizes to unseen target domains. Muandet et al. propose a kernel-based optimization algorithm to learn an invariant transformation. Li et al. introduce adversarial auto-encoders to learn a generalized latent feature representation across domains. Their maximum mean discrepancy measure aligns distributions to learn universal representations to be independent of domains. We explore the domain discrepancy to learn invariant representations through the lens of mutual information .
Information bottleneck (IB) provides an information-theoretic principle of encoding the input data into a compressed representation that maximizes target prediction. This is achieved by minimizing the mutual information between the input variable and its latent representation , while maximizing the mutual information between the output variable and the latent representation . To be more precise, the IB principle is to maximize the objective function:
where is the hyperparameter that controls the size of the information bottleneck, and are the corresponding model parameters.
The IB principle has recently been introduced for theoretical understanding and analysis of deep neural networks . The authors optimize the networks with an iterative Blahut-Arimoto algorithm, which is infeasible in practical systems. Alemi et al. developed a variational approximation to the IB objective by leveraging variational inference, which allows the IB model to be parameterized with neural networks. Amjad et al. investigated training deep neural networks (DNN) for classification based on minimization of the IB functional. It is shown that for deterministic DNNs, the optimization can be ill-posed. This is because the IB functional can be infinite or not admitting gradient descent since it is piece-wise constant. The possible remedy indicated in their work is to train stochastic DNNs with the IB principle.
Meta-learning, or learning to learn, endows models with the capacity to efficiently learn new tasks by acquiring common knowledge through experiencing a set of related tasks. It has been explored in several directions, e.g., by learning a meta learner on diverse tasks to adapt the parameters of the base learner on a specific task , learning to optimize the parameters of deep neural networks , and learning to learn the gradient optimization process by recurrent neural networks , etc. A representative meta-learning algorithm is the model agnostic meta-learning (MAML), which learns the models to be able to adapt to similar tasks with only a few gradient descent updates. Li et al. introduced the idea of MAML to domain generalization. They train models with generalization ability to unseen domains by leveraging the meta-learning setting. MetaReg addresses the domain shifts by leveraging the insights from meta-learning . They learn a meta regularizer to achieve the generalization from source to unseen target domains. Li et al. proposed a meta-learning approach based on a feature-critic network, in which an auxiliary loss is introduced to improve generalization ability. Dou et al. adopt a gradient-based model-agnostic learning algorithm to deal with domain shift for domain generalization. Two complementary losses are introduced for regularization of semantic features. The success of those works has indicated the effectiveness of meta-learning in domain generalization. Probabilistic meta-learning has also been developed in few-shot learning to handle uncertainty , which has not been explored for domain generalization.
In this work, we introduce a probabilistic meta-learning model for domain generalization which enables better handling of prediction uncertainty on unseen domains. We introduce the IB principle for domain-invariant representation learning by a stochastic deep neural network. We derive a new variational approximation to the IB principle under the meta-learning framework, resulting in the meta variational information bottleneck (MetaVIB) principle for domain generalization. We adapt the episodic training strategy in meta-learning by using the meta-train and meta-test splits of the source domains in each mini-batch for stochastic optimization.
Method
We describe the meta-learning setting for domain generalization. Following the setting in recent domain generalization by meta-learning , we divide a dataset into the Source domains used for train and the Target domains held-out for test. In the train phase, data in the source domains is episodically divided into sets of meta-train and meta-test domains. We train the model by optimizing over the prediction errors on meta-test domains. In the test phase, the learned model is applied to the target domains for performance evaluation. The training phase incorporates the idea of meta-learning which induces a higher level of learning by the split of meta-train and mete-test domains, rather than training on all source domains . This episodic meta-learning process mimics the generalization from seen to previously unseen domains.
We start with the probabilistic formulation of the domain generalization, based on which we develop the probabilistic model under the meta-learning framework. We consider the general estimation problem of conditionally predictive likelihood in the meta-test domain :
where is the sample of paired input and label drawn from data distribution in meta-test domain, is the conditionally predictive distribution, is the parameter set of the classifier. Note that we treat as a stochastic variable that depends on the input and the optimization of (2) is with respect to the parameters of probabilities.
In this work, we parameterize the model by deep neural networks. From the information-theoretic point of view , we regard the feature representation from the neural network as a stochastic variable , which is the latent encoding of the input . In domain generalization, it is commonly assumed that the label space is shared across the source and target domains. By leveraging the meta-learning setting, we propose to use data from the meta-train domains to estimate the parameters of the classifier by replacing with , which is applied to the meta-test domain. By incorporating the latent variable z into (2), we obtain the following maximum conditionally predictive likelihood estimation,
This establishes a probabilistic latent model which can be represented in a computational graph as shown in Fig. 1, and the corresponding conditional joint distribution is defined as:
where denotes the model parameters, , , and () are the number of samples in the meta-test (meta-train) domains. It is possible to directly employ (3) as the optimization objective using the techniques of amortized inference . However, the learned representations would not be domain invariant, which is desired for domain generalization. To achieve domain-invariant representations, we resort to the information bottleneck (IB) principle , which will be incorporated into the objective as a regularizer for joint optimization.
2 Meta Variational Information Bottleneck
We introduce the IB principle to learn domain-invariant representations under the meta-learning framework. We impose the information bottleneck on the feature representations to control the information flow in deep neural networks. This should largely remove domain related information while letting through the information that maximizes prediction of labels on the meta-test domain.
Let the random variables , , and denote the input, output, and the intermediate feature representation in the deep neural network, which encodes . The mutual information between the latent encoding of data and its output label is defined as follows:
Since is intractable, we introduce to be a variational approximation of , where conditioning on the classifier parameter is indicated by (4), and the prior distribution of is denoted as . Then we have:
where is the entropy of . Taking expectation values of both sides with respect to , we have
which is tractable in general by approximation .
Now we consider the second term , which can be written as follows:
By combining the two bounds (9) and (11), we establish the meta variational information bottleneck (MetaVIB)
which extends the IB theory into the meta-learning scenario, offering a new principle of learning domain-invariant representations for domain generalization.
We follow to approximate and with empirical data distribution and , where is the number of samples in the meta-test domain. This essentially regards the data points and as the samples drawn from the data distributions and , respectively.
We use Monte Carlo sampling to draw samples from for and from for in the lower bound of MetaVIB in (13). We attain the following objective function:
where is the number of classes and contains the samples from the -th category in the meta-train domains. We amortize the posterior distribution and the meta prior across classes, that is, the variational distribution of each class is inferred individually by the samples from its corresponding class , which further alleviates the computational overhead. In addition, the KL term can be calculated in a closed form. Here, to enable back-propagation, we adopt the re-parameterization trick , that is,
where is a deterministic function which is usually parameterized by a multiple layer perception (MLP) and and are the number of samples for and , respectively.
Taking a closer look at the objective (14), we observe that the first term is the negative log predictive likelihood in the meta-test domain, where the label of is predicted from its latent encoding and the classifier parameter . Minimizing the first term guarantees maximal prediction accuracy. The second term is the KL divergence between distributions of latent encoding of the sample in the target domain and that estimated by the samples from the same category in the meta-train domains. It is the minimization of the KL term in (14) that enables the model to learn domain-invariant representations. This is in contrast to the regular IB principle which is to compress the input and does not necessarily result in domain-invariant representations.
3 Learning with Stochastic Neural Networks
We implement the proposed model by end-to-end learning with stochastic neural networks that are comprised of convolutional layers and fully-connected layers. The inference is parameterized by a feed-forward multiple layer perception (MLP). During the training phase, given domains, we randomly sample one domain as the meta-test domain, the remaining domains are used as the meta-train domains. Then we choose a batch of samples from the meta-train domain , and a batch of samples from the meta-test domain . Note that samples from meta-train domains cover all the classes. For each sample of the -th class, we first extract its features via , where is the feature extraction network and we use permutation-invariant instance-pooling operations to get the mean feature of samples in the -th class. The mean feature will be fed into a small MLP network to calculate the mean and variance of the weight vector distribution for -th class, which is then used to sample the weight vector of this class by . The weight vectors of all classes are combined column by column to form a weight matrix .
We calculate the parameters of the latent distribution, i.e., the mean and variance of the -th class in the meta-train domain by another small MLP network . Then the parameter is sampled from the distribution . For each sample in the meta-test domain, we also calculate the mean and variance , of the distribution. Thus its latent coding vector can be naturally sampled from . Denote as the mean feature of all the samples of the -th class from the meta-train domains, i.e., . We provide the detailed step-by-step algorithm of the proposed MetaVIB for training in the supplemental material.
Experiments
We conduct our experiments on three benchmarks commonly used in domain generalization . We first provide ablation studies to gain insights into the properties and benefits of MetaVIB. Then we compare with previous methods based on both regular learning and meta-learning for domain generalization. We put more results in the supplementary material due to space limit.
VLCS is a real-world dataset that contains four domains collected from VOC2007 , LabelMe , Caltech-101 , and SUN09 . Images are from 5 classes, i.e., bird, car, chair, dog, person. The domain shift across those datasets makes VLCS a suitable benchmark for domain generalization.
PACS contains 9991 images from 4 domains, i.e., Photo, Art painting, Cartoon, and Sketch, which cover huge domain gaps. Images are from 7 object classes, i.e., dog, elephant, giraffe, guitar, horse, house, and person.
Rotated MNIST is a synthetic dataset consisting of 6 domains, each containing 1000 images of the 10 digits (i.e., , 100 for each) randomly selected from the training set of MNIST , with 6 rotation degrees: , and .
2 Implementation Details
Splits, Metrics and Backbone On all datasets, we follow the train-test splits suggested by , and perform experiments with the “leave-one-domain-out” strategy: we take the samples from one domain as the target domain for testing, and the samples from the remaining domains as the source domain for training. We use the AlexNet pre-trained on ImageNet and fine-tuned on the source domains of each dataset to perform testing on the target domain of that dataset. We use the average accuracy of all classes as the evaluation metric . To benchmark previous methods, we employ the pre-trained AlexNet on ImageNet as the backbone on VLCS and PACS. For Rotated MNIST we use a backbone network with two convolutions and one fully-connected layer. Even more implementation details about training stage, the feature extraction network and inference networks for different datasets are provided in the supplemental materials.
3 Ablation Study
To study the benefit of the MetaVIB under the probabilistic framework for domain generalization, we compare with several alternative models on VLCS and PACS in Tables 1 and 2.
To show the benefit of probabilistic modeling, we first consider AlexNet which is pre-trained on ImageNet, fine-tuned on the source domains and applied to the target domains. We define our Baseline model as the probabilistic model that predicts parameter distributions of the classifiers, without regular VIB or MetaVIB. The probabilistic model outperforms the pre-trained AlexNet by and on the VLCS and PACS benchmarks. The results indicate that the classifiers learned by probabilistic modeling better generalize to the target domains. The further analysis of the prediction uncertainty of the probabilistic modeling is put in the supplemental materials.
3.2 Benefit of MetaVIB
We show the benefit of MetaVIB by comparing with the regular VIB , which is applied to the baseline model as a regularization in the optimization, and the Baseline model. We first establish the probabilistic model with the regular VIB which performs better than the baseline (74.01% - up 0.64%) on VLCS and (73.37% - up 0.74%) on PACS. The VIB regularization term maximizes the mutual information between and the target , which will encourage better prediction performance compared to the Baseline model. However, our MetaVIB learns an even better domain-invariant representation, as it consistently outperforms VIB by up to on PACS . As indicated in the optimization objective in (14) minimizing the KL term makes the representations of samples in the meta-target domain to be close to the representations obtained by the samples of the same class from the meta-source domains. As a result, the learned model acquires the ability to generate domain-invariant representations by the episodic training. In contrast, the regular VIB is to simply compress the input with no explicit mechanism to narrow the gaps across domains. The obtained representations with regular VIB are not necessarily domain-invariant. Actually, there is no evident causal relation between compression and generalization as indicated in .
3.3 Influence of information bottleneck size β𝛽\beta
The bottleneck size controls the amount of information flow that goes through the bottleneck of the networks. To measure its influence on the performance, we plot the information plane dynamics of different network layers with varying in Fig. 2. We observe that MetaVIB with achieves the highest while at the same time is minimal. We also report the influence of in Table 3. MetaVIB achieves the best performance when , which is consistent with the information dynamic in Fig. 2. We observe in Fig. 2 (c) that with , the is lowest and is the highest, compared to those with other values of . A larger indicates that we can make more accurate predictions from , while a smaller indicates contains the minimal information from that is required for prediction, suggesting a domain-invariant representation . This explains why produces the best prediction results compared to other values of . In our experiments, the optimal value of is obtained by using a validation set for each dataset and we found produces the best performance on all datasets.
3.4 Analyzing domain-invariance
We visualize the features learned by the pre-trained AlexNet, VIB and MetaVIB in Fig. 3. For better illustration, we use t-SNE to reduce the feature dimension into a two-dimensional subspace. We observe that the features of the same category learned by pre-trained Alexnet (Fig. 3 (a)) show large discrepancy among the four domains. The regular VIB reduces this discrepancy to some extent, but still suffers from considerable gaps between the unseen domain (violet shapes) (Fig. 3 (b)). MetaVIB largely reduces the discrepancy of different domains including the unseen domains as shown in Fig. 3 (c). In Fig. 3 (d), we observe again that the gaps of features among 4 domains by the pre-trained AlexNet are larger than those between the 7 classes in each domain. Fig. 3 (e) shows that the VIB reduces the domain gaps to certain extent. From Fig. 3 (f), we observe MetaVIB reduces domain gaps considerably while at the same time scatters the samples of 7 classes in each domain. Overall, the proposed MetaVIB principle demonstrates effectiveness in learning domain-invariant representations to tackle domain shift.
3.5 Success and failure cases
We show some success and failure cases in Fig. 0.G.2. MetaVIB successfully predicts the labels for ambiguous images. The dog in the second image in Fig. 0.G.2 (a) wears human clothes, showing strong characteristics of a person. Yet, MetaVIB correctly predicts it with a high confidence probability of . The sketch of the horse looks like a dog in the fourth image, but MetaVIB predicts it correctly with a high probability of . In the failure cases (b), MetaVIB fails to make the correct prediction, but provides reasonable probabilities for both a person and a dog, which shows the effectiveness in handling uncertainty. It is hard to distinguish which object needs to be predicted in these images, as shown in the first image in Fig. 0.G.2 (b).
4 State-of-the-Art Comparison
We compare with regular and meta-learning methods for domain generalization. The results on the three datasets are reported in Tables 4-6. On the VLCS dataset , our MetaVIB achieves high recognition accuracy, surpassing the second best method, i.e., MASF , by a margin of . Note that on all domains, our MetaVIB consistently outperforms MLDG , which is a gradient-based meta-learning algorithm. On the PACS dataset , our MetaVIB again achieves the best overall performance. It outperforms most of the previous methods, showing clear performance advantages over JiGen . Again, our MetaVIB performs better than other meta-learning based methods, e.g., MetaReg , Reptile , MLDG , Feature-Critic , and MASF . It is worth highlighting that our MetaVIB exceeds those meta-learning methods on the “Cartoon” domain by phenomenal margins. On the Rotated MNIST dataset , the proposed MetaVIB achieves consistently high performance on the test domains, exceeding the alternative methods. It is worthwhile to mention that our MetaVIB outperforms the meta-learning algorithms MetaReg , and Reptile . showing its effectiveness as a meta-learning method for domain generalization. To conclude, on all datasets, our MetaVIB accomplishes better performance than previous methods based on both regular learning and meta-learning. The best results on all benchmarks validate the effectiveness of our method for domain generalization.
Conclusion
In this work, we propose a new probabilistic model for domain generalization under the meta-learning framework. To address prediction uncertainty, we model the parameters of the classifiers shared across domains by a probabilistic distribution, which is inferred from the source domain and directly used for the target domains. To reduce domain shift, our method learns domain-invariant representations by a new Meta Variational Information Bottleneck principle, derived from a variational bound of mutual information. MetaVIB integrates the strengths of meta-learning, variational inference and probabilistic modeling for domain generalization. Our MetaVIB has been evaluated by extensive experiments on three benchmark datasets for cross-domain visual recognition. Ablation studies validate the benefits of our contributions. MetaVIB consistently achieves high performance and advances the state of the art on all three benchmarks.
References
Appendix 0.A Algorithms of MetaVIB for Training
We describe the detailed algorithm for training MetaVIB as following Algorithm 1:
Appendix 0.B Learning Architecture
To better clearly understand our proposed MetaVIB, we draw a concise architecture diagram in Fig. 0.B.1.
Appendix 0.C Training Details
During the training, we use the Adam optimizer, and set the learning rate as . In each training batch, we randomly select three domains including two meta-train domains and one meta-test domain. In each domain, we choose samples, and the batch size is . The iteration number is set as . The model with the highest validation accuracy is employed to evaluate the test set from the meta-test domain.
Appendix 0.D Influence of information bottleneck size β𝛽\beta
We report Influence of information bottleneck size on the VLCS and Rotated MNIST in Tables 0.D.1 and 0.D.2. For the VLCS, MetaVIB obtains best results for , while for the Rotated MNIST, MetaVIB gets best results for .
Appendix 0.E Influence of the number of Monte Carlo Influence of the number of Monte Carlo samples
We use Monte Carlo sampling to draw samples from for . We report varying sample number on PACS in the Table 0.E.3. Our method achieves inferior results with ; performs consistently better with , converges at and becomes worse when . So in our experiments, we set and we averaged over runs on the test domain. The variance reflects the error caused by Monte Carlo sampling in each test experiment.
Appendix 0.F Network Architectures
The feature extraction network for PACS, VLCS is shown in Table 0.F.4, the feature extraction network for Rotated MNIST is shown in Table 0.F.5.
F.2 Inference Network
The architecture of the inference network for PACS, VLCS is in Table 0.F.6, the architecture of the inference network for Rotated MNIST is in Table 0.F.7.
The architecture of the inference network for PACS, VLCS is in Table 0.F.8, the architecture of the inference network for Rotated MNIST is in Table 0.F.9.
Appendix 0.G Prediction Uncertainty Analysis
Since the data follows distinct distribution between seen and unseen domains, uncertainty is inevitable during the prediction stage on the unseen domains, to which no data is accessible in the learning stage. To deal with the prediction uncertainty, we model parameters of classifiers shared across domains as probabilistic distributions that we infer from the data of the seen domains. The probabilistic modeling enables us to better handle the prediction uncertainty on previously unseen domains.
In order to demonstrate that the proposed probabilistic modeling can handle prediction uncertainty, we conduct an extra set of experiments as follows:
We shown more success and failure cases in Fig. 0.G.2 and show the corresponding prediction probabilities of using different sampled classifiers for each category of the image in Fig. 0.G.3-0.G.10. _ indicates the mean value of the classifier. From Fig. 0.G.3-0.G.10, we can see that different can produce different prediction probabilities to each category. Specially, for the fourth image of success cases, the final result of the classification is giraffe. However, the classifiers _ and _, our model predicts a higher prediction probability of horse than giraffe as shown in Fig. 0.G.6. For the fourth image of failure cases, the image is classified as dog, but that the prediction probability of elephant is higher than that of dog by using classifiers _ as shown in Fig. 0.G.10. Although the final prediction result of our model is incorrect, some of sampled classifiers can still make correct predictions.