Feature-Critic Networks for Heterogeneous Domain Generalization
Yiying Li, Yongxin Yang, Wei Zhou, Timothy M. Hospedales
Introduction
A shift in data statistics between training and testing is often unavoidable in real-world applications, and leads to a significant negative impact on the performance of machine learning models in practice. This motivates research into methods to ameliorate the impact of domain shift, including Domain Adaption (DA) (Bousmalis et al., 2016; Ganin & Lempitsky, 2015; Long et al., 2015, 2016) and Domain Generalisation (DG) (Muandet et al., 2013; Ghifary et al., 2015; Li et al., 2018a; Shankar et al., 2018).
Unsupervised Domain Adaptation (UDA) (Long et al., 2016; Saito et al., 2017; Shu et al., 2018) methods operate in the setting where we can access unlabelled testing (target) domain data during training to drive model adaptation and compensate for the domain shift. Domain Generalisation addresses the harder setting, where a model trained on a set of source domains should perform well on a novel target domain with different data statistics, without requiring any access to target domain data during training. That is, the model should be robust enough out-of-the-box to perform well in a new domain, without further parameter updates. Both DA and DG methods almost always assume the label space is consistent across both source and target domains.
In the case of disjoint label spaces between source and target domain, we term the domain generalisation problem as one of heterogeneous domain generalisation. In this case a feature representation trained on a source domain should generalise to supporting recognition of novel categories in a novel target domain. This problem setting is actually widely encountered. The central example is the ubiquitous computer vision pipeline where a CNN feature extractor pre-trained on ImageNet is re-used for diverse applications. If data, computation, and human expert time is available, the feature can be fine-tuned on the target problem. However, for many practical applications lacking one or more of these requirements, standard practice is to use an ImageNet CNN off-the-shelf as a fixed feature extractor, and train a shallow model such as SVM or KNN for the new problem (Donahue et al., 2014; Razavian et al., 2014). This pipeline is an example of the heterogeneous domain generalisation setting, in that a feature is being asked to generalise to supporting recognition of novel categories in data with novel statistics. The ImageNet pre-trained feature is strong enough to do a reasonable job of this already. However, given the ubiquity of this pipeline, providing an improved general purpose feature would be widely beneficial. In this paper we aim to do exactly this by presenting a novel method that explicitly trains a feature to prepare it for domain and label shift. We demonstrate this via performing heterogeneous DG on the Visual Decathalon benchmark (Rebuffi et al., 2017). This also provides the largest scale evaluation of DG to date.
We are inspired by recent meta-learning learning methods that perform episodic training (Finn et al., 2017; Snell et al., 2017; Ravi & Larochelle, 2017) to simulate the train/test process to improve few-shot learning. In this work, we propose to perform meta-learning to improve feature extractor training, and deliver a better model for both homogeneous and heterogeneous DG problems.
To realise our idea, we simulate training-to-testing domain shift by splitting our source domains into virtual training and testing (i.e., validation) domains. The source model is decomposed into feature extractor and task networks (i.e., a classifier network in our case). Crucially we then introduce a feature-critic network that learns to criticise the quality of the features produced by the feature network, specifically with regards to their robustness to the simulated domain shift. This feature-critic provides a learned auxiliary loss which provides an additional source of feedback to the feature network (besides the conventional supervised classification loss via the task network), and enables it to produce a more robust feature. The feature, task and critic networks are trained together end-to-end in a meta-learning pipeline. Our evaluation shows good performance in the conventional DG setting using Rotated MNIST (Ghifary et al., 2015; Motiian et al., 2017) and PACS (Li et al., 2017a) benchmarks, as well as the heterogeneous DG setting using the larger scale Visual Decathlon (VD) (Rebuffi et al., 2017) benchmark.
Related Work
Multi-Domain Learning (MDL) MDL addresses training a single model capable of solving multiple datasets (domains). If the data is relatively small and the domains are similar, this sharing can lead to improved performance compared to training a separate model per domain (Yang & Hospedales, 2015). On the other hand, for diverse domains with large data, MDL may under-perform a single model per domain; but is nonetheless is of interest due to the simplicity of a single model and its better memory scalability compared to a separate model per domain (Rebuffi et al., 2017, 2018). We mention MDL here, because DG methods typically train on multiple source domains as per MDL – but furthermore aim to generalise to novel held out domains.
Domain Generalisation (DG) DG relates to domain-adaptation in that we care about performance on a target domain, rather than source domains; however it considers the case where target domain samples are unavailable during training, so the model must generalise directly rather than adapt to the target domain. DG is of related to conventional generalisation: where models learned on a set of training instances generalise to novel testing instances, for example by regularisation. However it operates at a higher level, where we aim to help models trained on a set of training domains generalise to a novel testing domain.
Most existing DG approaches can be split into three categories: feature-based methods, classifier-based methods, and data augmentation methods. Feature-based methods: These aim to generate a domain-invariant representation. For example where the distance between the empirical distributions of the source and target examples is minimized (Li et al., 2018b; Muandet et al., 2013; Li et al., 2018a). Classifier-based methods: Some aim to enhance generalisation by fusing multiple sub-classifiers learned from source domains (Duan et al., 2012; Niu et al., 2015a, b), and others learn an improved classifier regulariser using source samples – notably the recently proposed MetaReg (Balaji et al., 2018). Data augmentation methods: CrossGrad (Shankar et al., 2018) generates provides domain-guided perturbations of input instances, which are then used to train a more robust model. Volpi et al. (2018) defines an adaptive data augmentation scheme by appending adversarial examples at each iteration. Our Feature-Critic approach falls into the feature-based category, but meta-learns a feature-critic network to train a robust shared feature extractor.
Few studies have considered the heterogeneous DG setting, where the domains do not share the same label space. We do not expect the classifier to generalise directly to the target domain (impossible due to the change in label space), but we do aim to improve the robustness of a source-domain trained feature in terms of its generalisation to successfully represent a novel problem. Most existing DG methods cannot be applied here. We show how to modify MetaReg (Balaji et al., 2018) and Reptile (Nichol et al., 2018) to address this DG setting. The most relevant benchmark is Visual Decathlon (VD) (Rebuffi et al., 2017). The VD benchmark was proposed to evaluate multi-domain and lifelong (Rosenfeld & Tsotsos, 2018) learning. We re-purpose it for DG evaluation. In this case a model trained on the six largest datasets in VD should produce a feature which provides a general and robust enough encoding to allow the four smaller datasets to be classified with a simple shallow classifier.
Meta-Learning Meta-learning (a.k.a. learning to learn, (Schmidhuber et al., 1997; Thrun & Pratt, 1998)) has received resurgence in interest recently with applications in few-shot learning (Li et al., 2017b; Snell et al., 2017; Sung et al., 2018) and beyond (Xu et al., 2018). In few-shot meta-learning, a common strategy is to simulate the few-shot learning scenario by randomly drawing few-shot train/test episodes from the full training set. We adapt this episodic training strategy by creating virtual training and testing splits of our source domains in each mini-batch.
A few methods have applied related episodic meta-learning strategies in DG (Li et al., 2018a; Balaji et al., 2018). MLDG (Li et al., 2018a) defined a heuristic gradient descent update rule based on the gradients of the simulated training and testing domains. MetaReg (Balaji et al., 2018) trains the weights of the classifier’s regulariser so as to produce a more general classifier for a fixed feature extractor. In contrast, our Feature-Critic produces a more general feature extractor that can be used with any classifier. This is achieved by simultaneously learning an auxiliary loss function (Gygli et al., 2017; Sung et al., 2017) (i.e., the critic network) that trains the feature extractor for improved domain invariance.
Methodology
We introduce the proposed method under the heterogeneous DG setting, but it is straightforwardly applicable to conventional (homogeneous) DG as a special case. Assuming that we have N domains (datasets) , and each domain contains a set of data-label pairs, i.e., . We also have the training split of target (testing) domain, , but we can not access this for feature learning.
We assume a CNN model split into two parts: feature extractor and classifier . For heterogeneous DG, we have classifiers, denoted , and a universal feature extractor shared for all domains (assuming that images from all domains are resized to the same size). In the homogeneous DG, we only need a single classifier that can be shared across all domains.
A naive deep learning approach called aggregation (AGG) trains a single extractor to minimise the total cross-entropy (CE) loss of all domains.
Here is the th domain and is a mini-batch of it.
This simple baseline surpasses many prior purpose designed DG methods as discussed in Li et al. (2017a). The key question is how to improve this naive approach, such that the trained feature extractor produces more robust features that generalise better to unseen target domains.
2 Simulating Domain Shift in Training
Our high-level strategy simulates domain-shift during training as illustrated in Algo. 1 and Figure 1. We use the learned feature-critic loss to guide learning on the meta-training set , and optimise the feature-critic itself on the meta-validation set . The key idea is that training with on should improve its performance on .
3 Meta-Learning an Auxiliary loss
Thus we train the auxiliary loss (feature-critic network) to promote this. Specifically, we optimise the parameter of feature-critic network as follows:
Here is a function that measures the validation domain performance (larger is better), and we discuss how to design it in Sec. 3.4. is a utility function, which converts the reward (performance gain) to utility. It reflects the commonly accepted idea concept diminishing marginal utility, and links with . If and the term are excluded in Eq. 3, it would simply maximise the validation set performance with . The reason Eq. 3 is better is that serves as a baseline, making the value range – and thus the gradient – more stable. One can understand the role of here as a smoother version of min/max-margin or a softer version of gradient clipping.
4 Measuring Validation Performance: Designing γ𝛾\gamma
To measure validation performance, can take up to four variables as input: feature extractor parameter , classifier parameter , data , and label , i.e., . One simple choice is the negative classification loss, i.e.,
Therefore rather than designing to take directly, we propose a more efficient and effective way to enable to promote the base network’s generalisation. Specifically, the auxiliary loss operates on the extracted features . Since our auxiliary generalisation-promoting loss operates on the feature representation produced by the base network, we denote it Feature-Critic.
Denote as the sized matrix stacking the -dimensional features from examples in a mini-batch from the th domain in the virtual training set . . A key requirement of is that it should be permutation invariant to the rows of , i.e., it should not make a difference if we feed images indexed . Two available choices are:
(i) The set embedding (Zaheer et al., 2017), i.e.,
where denotes a row of , and is a multi-layer perceptron.
(ii) The flattened covariance matrix, i.e.,
Finally, the MLP’s output should be a scalar and we place a softplus activation to make sure its output is non-negative.
6 Summary
Bringing all the components together, we have the full Algo. 2. To summarise, we randomly draw train/validation domains in each iteration and: Perform a putative feature extractor update on with and without the auxiliary Feature-Critic loss. Then generate a meta-loss based on whether or not the feature extractor update has improved performance on the validation set. Finally the feature extractor/classifier are updated using the supervised and auxiliary losses, and auxiliary loss itself is updated using the meta-loss.
Experiments
We evaluate our approach, first on the heterogeneous DG problem using the VD benchmark (Section 4.1), and then on the conventional homogeneous DG using Rotated MNIST and PACS (Section 4.2). Our demo code can be viewed on https://github.com/liyiying/Feature_Critic.
Dataset The Visual Decathlon dataset, initially proposed for multi-domain learning (Rebuffi et al., 2017), also provides a large scale and rigorous benchmark for DG. VD contains ten diverse domains including handwritten characters, pedestrians, traffic signs, etc. The images have been pre-processed to . To use this benchmark for DG, we aim to train a network on a subset of source domains, and produce a robust feature extractor that provides a good representation for classification in a disjoint subset of target domains. It should do so ‘out-of-the-box’, without further fine tuning. Specifically, we take the six larger datasets (CIFAR-100, Daimler Ped, GTSRB, Omniglot, SVHN and ImageNet) as source domains and hold out the four smaller datasets (Aircraft, D. Textures, VGG-Flowers and UCF101) as target domains. We use ImageNet pre-trained ResNet-18 (He et al., 2016) as the base network for all competitors. For computational efficiency, we freeze the first four blocks of ResNet-18 and only update the remaining blocks, as well as the average pooling layer, during DG training. For all methods, the final feature is used to train SVM or KNN for the target task. All methods are evaluated by both average multi-class classification accuracy in the target domains, as well as the VD-Score metric (Rebuffi et al., 2017) that rewards consistently high performance across all domains.
Competitors Few competitors can address heterogeneous DG. For these we consider AGG baseline (Eq 1), CrossGrad (Shankar et al., 2018), MetaReg (Balaji et al., 2018), and Reptile (Nichol et al., 2018). MetaReg is originally designed to produce a robust classifier given a fixed feature. We modify MetaReg to support the heterogeneous DG by (i) applying it on the feature extractor instead (as per our Feature-Critic), called MR; (ii) applying it on the final layer of feature extractor, called MR-FL. Meanwhile Reptile is designed for few-shot meta-learning. However after modifying it for multi-domain rather than multi-task meta-learning, we found it effective for heterogeneous DG.
Feature-Critic Settings We use the set embedding architecture for the critic network (Eq 6), as the covariance architecture requires too many parameters using high dimensional ResNet. During each iteration, we randomly choose four of the six source domains as meta-train, and the remaining two provide the meta-test (validation) domains. We train all components end-to-end using the AMSGrad (Reddi et al., 2018) (batch-size/per meta-train domain=64, batch-size/per meta-test domain=32, lr=0.0005, weight decay=0.0001) for 30k iterations where the lr decayed in 5K, 12K, 15K, 20K iterations by a factor 5, 10, 50, 100, respectively. Similar to MetaReg (Balaji et al., 2018), after the parameters are trained via meta-learning, we fine-tune the network on all source datasets for the final 10k iterations.
Results We first assume the full training split is available for each target domain. Table 1 shows that: (i) The original ImageNet feature transfers to novel tasks reasonably well, as observed by classic studies (Yosinski et al., 2014). (ii) Demonstrating the benefit of simply exploiting large datasets, the AGG baseline’s feature, trained on more than 1.39 million images across the six domains, provides strong performance. However, while it has a higher average accuracy than the ImageNet feature, AGG’s VD score is lower, reflecting its inconsistent performance. Thus obtaining consistently high scores from multi-domain training is non-trivial. Naively aggregating more diverse source domains into training can both help and hinder performance (for example, depending on if aggregated domains are particularly similar or dissimilar to a given target). Nevertheless, AGG sometimes outperforms prior purpose designed DG methods CrossGrad and MetaReg, with only Reptile producing a feature that outperforms AGG in both accuracy and VD-score metrics. (iii) Overall, our Feature-Critic (FC) method generally provides the best performance across domains and across both types of classifiers evaluated.
Although the above application scenario of heterogeneous DG is one where compute, memory or human resources rule out feature fine-tuning, another motivating scenario is where the target domain data is too sparse for effective fine-tuning. Thus we next investigate the situation if less target data is available. Specifically, we repeat the evaluation assuming that [10, 25, 50, 100] of the training split is available for SVM/KNN training. Table 2 reports target domain test accuracies under these settings. We can see that Feature-Critic provides a consistent improvement over the alternatives. Finally, we also consider a genuinely few-shot setting for the target domain. In this case we consider labelled examples per class in the target domain, and perform KNN recognition on their test sets. The results in Table 3 show that for simple similarity-based matching in a novel target domain, Feature-Critic also provides the best off-the-shelf feature representation.
In summary the Feature-Critic meta-training strategy produces a feature extractor that is generally useful for diverse target problems in an off-the-shelf feature + shallow classifier configuration. The results outperform both the standard ImageNet feature and the obvious Data Aggregation extension across a range of operating points in the target domain from the few to many-shot regime. This suggests that Feature-Critic trained feature extractors are of wide potential value in diverse applications.
2 Homogeneous DG experiments
Dataset and Settings Rotated MNIST (Ghifary et al., 2015) contains six domains with each corresponding to a degree of roll rotation in the classic MNIST dataset. The basic view (M0) is formed by randomly choosing 100 images each of ten classes from the original MNIST and we create 5 rotating domains from M0 with rotation each in clockwise direction, denoted by M15, M30, M45, M60, and M75. Following the setting in (Shankar et al., 2018), we perform leave-one-domain-out experiments by picking one domain to hold out as the target. We compare AGG baseline, as well as CrossGrad and MetaReg. For a recognition network all competitors use the standard MNIST CNN with two conv and one FC layer as the feature network and another FC layer as the classifier. We note prior studies (Ghifary et al., 2015; Shankar et al., 2018; Deshmukh et al., 2017) did not release specific selection of digits from within MNIST, so our results do not match the numbers in those papers exactly. However, we repeat all experiments 10 times and report the mean and standard deviation of recognition accuracy.
For Feature-Critic, we train using the AMSGrad optimizer (lr=0.001, weight decay=0.00005) for 5,000 iterations. For each iteration, one meta-train and one meta-test domain are chosen randomly from the five source domains. We also use this opportunity to compare the two variants of our loss function: Feature-Critic-MLP and Feature-Critic-Flatten.
Results It can be seen from Table 4 that AGG is again a strong baseline to beat. Over ten trials of 1000 digit samples, CrossGrad and MetaReg failed to match AGG, with only Reptile matching AGG’s performance. Meanwhile, Feature-Critic performs well with both variants of the auxiliary loss network, with the set embedding (Eq. 6) performing slightly better than the covariance matrix embedding (Eq. 7).
To qualitatively visualise the results we perform PCA projections of the features in the target domain. Figure 2 shows these projections, taking as an example the M15 domain as held out. Each dot denotes an image and the colour denotes its label. We can see that Feature-Critic (take MLP style as an example) feature extractor provides improved separability in the target domain compared to the AGG baseline.
Further Analysis Figure 3 reports the loss curves of cross entropy loss, auxiliary loss, and meta-loss for Feature-Critic during training. The cross entropy loss converges to zero, as the network usually can fit the training data perfectly. The auxiliary loss fluctuates up and down for the early stage of training, and finally stabilises to a small value. Because the auxiliary loss function itself is learned, its behaviour changes with its own learning process, which explains the fluctuations, esp. for the early stage. It is more interesting to see the pattern of meta-loss, which is the performance difference of feature extractor parameterised by and that by . If we select zero as a threshold, meta-loss has a clear pattern: “above zero” “below zero” “’being zero’. This pattern is expected because: (i) For the early stage, the auxiliary loss’ parameters are randomly initialised, so it knows little about how to help generalise, thus the gradients produced by it are rather random and less likely to help. Thus -based model outperforms -based model. (ii) With the updating of , improves and begins to make better than . During this period, gradients produced by the auxiliary loss help the model learn to generalise. (iii) For the late stage, meta-loss goes towards zero, which indicates that no longer helps (but it does not hurt either), as all of its knowledge has now been distilled into the feature extractor. The pattern of the three losses also demonstrates that, empirically, the whole algorithm converges, including the learned auxiliary loss.
2.2 Evaluation on PACS dataset
Dataset and Settings PACS (Li et al., 2017a) is a recent object recognition benchmark for domain generalisation. PACS contains 9991 images of size from four different domains - Photo, Art painting, Cartoon and Sketch. It has 7 categories across these domains: dog, elephant, giraffe, guitar, house, horse and person. We follow the standard protocol and perform leave-one-domain-out evaluation. Beyond this there have been two splits of PACS used in the literature. PACS was defined with a train/validation/test split within each domain. In Li et al. (2018a) models are trained on the train split alone with the validation split used for early stopping. In Balaji et al. (2018) the combined train+validation splits were used to train the models, resulting in slightly higher performance due to more data. For direct comparison with previously published results we evaluate both of these settings.
The ImageNet pre-trained AlexNet (Krizhevsky et al., 2012) is used as the backbone network. Our competitors include: DICA (Muandet et al., 2013), D-MTAE (Ghifary et al., 2015), DSN (Bousmalis et al., 2016), TF-CNN (Li et al., 2017a), MLDG (Li et al., 2018a), DANN (Ganin et al., 2016), CIDDG (Li et al., 2018c), Reptile (Nichol et al., 2018), CrossGrad (Shankar et al., 2018) and MetaReg (Balaji et al., 2018). We note that DANN is designed for domain adaptation, and Reptile for few-shot learning. We re-purpose them for DG. DANN, Reptile, CrossGrad, AGG, and MetaReg in Table 5 are our implementations. The other results are taken from Li et al. (2018a), Li et al. (2018c) and Balaji et al. (2018). Among these, MetaReg makes a domain general classifier; MLDG aligns gradients to achieve a more robust optima; CrossGrad synthesises data for a new domain; DANN makes indistinguishable representations across source domains; CIDDG learns discriminative features to match the distributions across domains; Feature-Critic (FC) learns representations that generalise to new domains with a better objective since it trains a supervised loss and explicitly simulates domain shift.
Our Feature-Critic (set embedding variant) is trained with M-SGD optimizer (batch size/per meta-trian domain=32, batch size/per meta-test domain=16, lr=0.0005, weight decay=0.00005, momentum=0.9) for 45K iterations. At each iteration, we randomly choose two of the three source domains as meta-train and the remaining one as meta-test.
Results The comparison with state-of-the-art methods on PACS dataset is shown in Table 5 and Table 6. AGG provides a hard baseline to beat as usual. Nevertheless Feature-Critic performs comparably to the best performing state of the art alternative in both settings of this benchmark.
Conclusion
We addressed the domain generalisation problem with a particular focus on the heterogeneous case, by meta-learning a regulariser to help train a feature extractor to be domain invariant. The resulting feature extractor outperforms alternatives for general purpose use as a fixed downstream image encoding. Evaluated on Visual Decathlon – the largest DG evaluation thus far – this suggests that Feature-Critic trained feature extractors could be of wide potential value in diverse applications. Furthermore Feature-Critic also performs favourably compared to state-of-the-art in the homogeneous DG setting. In future work we will apply Feature-Critic to other problems including RL, and explore the impact on fine-tuning target problems.
Acknowledgements
This work was produced while the first author was visiting the University of Edinburgh. This work was supported by National Natural Science Foundation of China (Grant No. 61751208), China Scholarship Council, EPSRC grant EP/R026173/1, and NVIDIA Corporation GPU donation.