Visual Representation Learning Does Not Generalize Strongly Within the Same Domain
Lukas Schott, Julius von Kügelgen, Frederik Träuble, Peter Gehler, Chris Russell, Matthias Bethge, Bernhard Schölkopf, Francesco Locatello, Wieland Brendel
Introduction
Humans excel at learning underlying physical mechanisms or inner workings of a system from observations (Funke et al. 2021; Barrett et al. 2018; Santoro et al. 2017; Villalobos et al. 2020; Spelke 1990), which helps them generalize quickly to new situations and to learn efficiently from little data (Battaglia et al. 2013; Dehaene 2020; Lake et al. 2017; Téglás et al. 2011). In contrast, machine learning systems typically require large amounts of curated data and still mostly fail to generalize to out-of-distribution (OOD) scenarios (Schölkopf et al. 2021; Hendrycks & Dietterich 2019; Karahan et al. 2016; Michaelis et al. 2019; Roy et al. 2018; Azulay & Weiss 2019; Barbu et al. 2019). It has been hypothesized that this failure of machine learning systems is due to shortcut learning (Kilbertus* et al. 2018; Ilyas et al. 2019; Geirhos et al. 2020; Schölkopf et al. 2021). In essence, machines seemingly learn to solve the tasks they have been trained on using auxiliary and spurious statistical relationships in the data, rather than true mechanistic relationships. Pragmatically, models relying on statistical relationships tend to fail if tested outside their training distribution, while models relying on (approximately) the true underlying mechanisms tend to generalize well to novel scenarios (Barrett et al. 2018; Funke et al. 2021; Wu et al. 2019; Zhang et al. 2018; Parascandolo et al. 2018; Schölkopf et al. 2021; Locatello et al. 2020a; Locatello et al. 2020b). To learn effective statistical relationships, the training data needs to cover most combinations of factors of variation (like shape, size, color, viewpoint, etc.). Unfortunately, the number of combinations scales exponentially with the number of factors. In contrast, learning the underlying mechanisms behind the factors of variation should greatly reduce the need for training data and scale more gently with the number of factors (Schölkopf et al. 2021; Peters et al. 2017; Besserve et al. 2021).
Benchmark: Our goal is to quantify how well machine learning models already learn the mechanisms underlying a data generative process. To this end, we consider four image data sets where each image is described by a small number of independently controllable factors of variation such as scale, color, or size. We split the training and test data such that models that learned the underlying mechanisms should generalize to the test data. More precisely, we propose several systematic out-of-distribution (OOD) test splits like composition (e.g., train = small hearts, large squares → test = small squares, large hearts), interpolation (e.g., small hearts, large hearts → medium hearts) and extrapolation (e.g., small hearts, medium hearts → large hearts). While the factors of variation are independently controllable (e.g., there may exist large and small hearts), the observations may exhibit spurious statistical dependencies (e.g., observed hearts are typically small, but size may not be predictive at test time). Based on this setup, we benchmark 17 representation learning approaches and study their inductive biases. The considered approaches stem from un-/weakly supervised disentanglement, supervised learning, and the transfer learning literature.
Results: Our benchmark results indicate that the tested models mostly struggle to learn the underlying mechanisms regardless of supervision signal and architecture. As soon as a factor of variation is outside the training distribution, models consistently tend to predict a value in the previously observed range. On the other hand, these models can be fairly modular in the sense that predictions of in-distribution factors remain accurate, which is in part against common criticisms of deep neural networks (Greff et al. 2020; Csordás et al. 2021; Marcus 2018; Lake & Baroni 2018).
New Dataset: Previous datasets with independent controllable factors such as dSprites, Shapes3D, and MPI3D (Matthey et al. 2017; Kim & Mnih 2018; Gondal et al. 2019) stem from highly structured environments. For these datasets, common factors of variations are scaling, rotation and simple geometrical shapes. We introduce a dataset derived from celebrity faces, named CelebGlow, with factors of variations such as smiling, age and hair-color. It also contains all possible factor combinations. It is based on latent traversals of a pretrained Glow network provided by Kingma et al. (Kingma & Dhariwal 2018) and the Celeb-HQ dataset (Liu et al. 2015).
We hope that this benchmark can guide future efforts to find machine learning models capable of understanding the true underlying mechanisms in the data. To this end, all data sets and evaluation scripts are released alongside a leaderboard on GitHub. https://github.com/bethgelab/InDomainGeneralizationBenchmark
Problem setting
Inductive biases for generalization in visual representation learning
We now explore different types of assumptions, or inductive biases, on the representational format (section 3.1), architecture (section 3.2), and dataset (section 3.3) which have been proposed and used in the past to facilitate generalization. Inductive inference and the generalization of empirical findings is a fundamental problem of science that has a long-standing history in many disciplines. Notable examples include Occam’s razor, Solomonoff’s inductive inference (Solomonoff 1964), Kolmogorov complexity (Kolmogorov 1998), the bias-variance-tradeoff (Kohavi et al. 1996; Von Luxburg & Schölkopf 2011), and the no free lunch theorem (Wolpert 1996; Wolpert & Macready 1997). In the context of statistical learning, Vapnik and Chervonenkis (Vapnik & Chervonenkis 1982; Vapnik 1995) showed that generalizing from a sample to its population (i.e., IID generalization) requires restricting the capacity of the class of candidate functions—a type of inductive bias. Since shifts between train and test distributions violate the IID assumption, however, statistical learning theory does not directly apply to our types of OOD generalization.
OOD generalization across different (e.g., observational and experimental) conditions also bears connections to causal inference (Pearl 2009; Peters et al. 2017; Hernán & Robins 2020). Typically, a causal graph encodes assumptions about the relation between different distributions and is used to decide how to “transport” a learned model (Pearl & Bareinboim 2011; Pearl et al. 2014; Bareinboim & Pearl 2016; von Kügelgen et al. 2019). Other approaches aim to learn a model which leads to invariant prediction across multiple environments (Schölkopf et al. 2012; Peters et al. 2016; Heinze-Deml et al. 2018; Rojas-Carulla et al. 2018; Arjovsky et al. 2019; Lu et al. 2021). However, these methods either consider a small number of causally meaningful variables in combination with domain knowledge, or assume access to data from multiple environments. In our setting, on the other hand, we aim to learn from higher-dimensional observations and to generalize from a single training set to a different test environment.
Our work focuses on OOD generalization in the context of visual representation learning, where deep learning has excelled over traditional learning approaches (Krizhevsky et al. 2012; LeCun et al. 2015; Schmidhuber 2015; Goodfellow et al. 2016). In the following, we therefore concentrate on inductive biases specific to deep neural networks (Goyal & Bengio 2020) on visual data. For details regarding specific objective functions, architectures, and training, we refer to the supplement.
Learning useful representations of high-dimensional data is clearly important for the downstream performance of machine learning models (Bengio et al. 2013). The first type of inductive bias we consider is therefore the representational format. A common approach to representation learning is to postulate independent latent variables which give rise to the data, and try to infer these in an unsupervised fashion. This is the idea behind independent component analysis (ICA) (Comon 1994; Hyvärinen & Oja 2000) and has also been studied under the term disentanglement (Bengio et al. 2013). Most recent approaches learn a deep generative model based on the variational auto-encoder (VAE) framework (Kingma & Welling 2013; Rezende et al. 2014), typically by adding regularization terms to the objective which further encourage independence between latents (Higgins et al. 2017; Kim & Mnih 2018; Chen et al. 2018; Kumar et al. 2018; Burgess et al. 2018).
2 Inductive bias 2: architectural (supervised learning)
The physical world is governed by symmetries (Nother 1915), and enforcing appropriate task-dependent symmetries in our function class may facilitate more efficient learning and generalization. The second type of inductive bias we consider thus regards properties of the learned regression function, which we refer to as architectural bias. Of central importance are the concepts of invariance (changes in input should not lead to changes in output) and equivariance (changes in input should lead to proportional changes in output). In vision tasks, for example, object localization exhibits equivariance to translation, whereas object classification exhibits invariance to translation. E.g., translating an object in an input image should lead to an equal shift in the predicted bounding box (equivariance), but should not affect the predicted object class (invariance).
A famous example is the convolution operation which yields translation equivariance and forms the basis of convolutional neural networks (CNNs) (Le Cun et al. 1989; LeCun et al. 1989). Combined with a set operation such as pooling, CNNs then achieve translation invariance. More recently, the idea of building equivariance properties into neural architectures has also been successfully applied to more general transformations such as rotation and scale (Cohen & Welling 2016; Cohen et al. 2019; Weiler & Cesa 2019) or (coordinate) permutations (Zhang et al. 2019; Achlioptas et al. 2018). Other approaches consider affine transformations (Jaderberg et al. 2015), allow to trade off invariance vs dependence on coordinates (Liu et al. 2018), or use residual blocks and skip connections to promote feature re-use and facilitate more efficient gradient computation (He et al. 2016; Huang et al. 2017). While powerful in principle, a key challenge is that relevant equivariances for a given problem may be unknown a priori or hard to enforce architecturally. E.g., 3D rotational equivariance is not easily captured for 2D-projected images, as for the MPI3D data set.
In our study, we consider the following architectures: standard MLPs and CNNs, CoordConv (Liu et al. 2018) and coordinate-based (Sitzmann et al. 2020) nets, Rotationally-Equivariant (Rotation-EQ) CNNs (Cohen & Welling 2016), Spatial Transformers (STN) (Jaderberg et al. 2015), ResNet (RN) 50 and 101 (He et al. 2016), and DenseNet (Huang et al. 2017). All networks are trained to directly predict the FoVs in a purely supervised fashion.
3 Inductive bias 3: leveraging additional data (transfer learning)
The physical world is modular: many patterns and structures reoccur across a variety of settings. Thus, the third and final type of inductive bias we consider is leveraging additional data through transfer learning. Especially in vision, it has been found that low-level features such as edges or simple textures are consistently learned in the first layers of neural networks, which suggests their usefulness across a wide range of tasks (Sun et al. 2017). State-of-the-art approaches therefore often rely on pre-training on enormous image corpora prior to fine-tuning on data from the target task (Kolesnikov et al. 2020; Mahajan et al. 2018; Xie et al. 2020). The guiding intuition is that additional data helps to learn common features and symmetries and thus enables a more efficient use of the (typically small amount of) labeled training data. Leveraging additional data as an inductive bias is connected to the representational format section 3.1 as they are often combined during pre-training.
In our study, we consider three pre-trained models: RN-50 and RN-101 pretrained on ImageNet-21k (Deng et al. 2009; Kolesnikov et al. 2020) and a DenseNet pretrained on ImageNet-1k (ILSVRC) (Russakovsky et al. 2015). We replace the last layer with a randomly initialized readout layer chosen to match the dimension of the FoVs of a given dataset and fine-tune the whole network for 50,000 iterations on the respective train splits.
Experimental setup
We consider datasets with images generated from a set of discrete Factors of Variation (FoVs) following a deterministic generative model. All selected datasets are designed such that all possible combinations of factors of variation are realized in a corresponding image. dSprites (Matthey et al. 2017), is composed of low resolution binary images of basic shapes with 5 FoVs: shape, scale, orientation, x-position, and y-position. Next, Shapes3D (Kim & Mnih 2018), a popular dataset with 3D shapes in a room with 6 FoVs: floor, color, wall color, object color, object size, object type, and camera azimuth. Furthermore, with CelebGlow we introduce a novel dataset that has more natural factors of variations such as smiling, hair-color and age. For more details and samples, we refer to Appendix B. Lastly, we consider the challenging and realistic MPI3D (Gondal et al. 2019), which contains real images of physical 3D objects attached to a robotic finger generated with 7 FoVs: color, shape, size, height, background color, x-axis, and y-axis. For more details, we refer to Section H.1.
2 Splits
Composition: We exclude all images from the train split if factors are located in a particular corner of the FoV hyper cube given by all FoVs. This means certain systematic combinations of FoVs are never seen during training even though the value of each factor is individually present in the train set. The related test split then represents images of which at least two factors resemble such an unseen composition of factor values, thus testing generalization w.r.t. composition. Interpolation: Within the range of values of each FoV, we periodically exclude values from the train split. The corresponding test split then represents images of which at least one factor takes one of the unseen factor values in between, thus testing generalization w.r.t. interpolation. Extrapolation: We exclude all combinations having factors with values above a certain label threshold from the train split. The corresponding test split then represents images with one or more extrapolated factor values, thus testing generalization w.r.t. extrapolation. Random: Lastly, as a baseline to test our models performances across the full dataset in distribution, we cover the case of an IID sampled train and test set split from . Compared to inter- and extrapolation where factors are systematically excluded, here it is very likely that all individual factor values have been observed in a some combination.
3 Evaluation
To benchmark the generalization capabilities, we compute the -score, the coefficient of determination, on the respective test set. We define the -score based on the MSE score per FoV
where is the variance per factor defined on the full dataset . Under this score, can be interpreted as perfect regression and prediction under the respective test set whereas indicates random guessing with the MSE being identical to the variance per factor. For visualization purposes, we clip the to 0 if it is negative. We provide all unclipped values in the Appendix.
Experiments and results
Our goal is to investigate how different visual representation models perform on our proposed systematic out-of-distribution (OOD) test sets. We consider un-/weakly supervised, fully supervised, and transfer learning models. We focus our conclusions on MPI3D-Real as it is the most realistic dataset. Further results on dSprites and Shapes3D are, however, mostly consistent.
In the first subsection, Section 5.1, we investigate the overall model OOD performance. In Sections 5.2 and 5.3, we focus on a more in-depth error analysis by controlling the splits s.t. only a single factor is OOD during testing. Lastly, in Section 5.4, we investigate the connection between the degree of disentanglement and downstream performance.
In Figure 4 and Appendix Figure 11, we plot the performance of each model across different generalization settings. Compared to the in-distribution (ID) setting (random), we observe large drops in performance when evaluating our OOD test sets on all considered datasets. This effect is most prominent on MPI3D-Real. Here, we further see that, on average, the performances seem to increase as we increase the supervision signal (comparing RN50, RN101, DenseNet with and without additional data on MPI3D). On CelebGlow, models also struggle to extrapolate. However, the results on composition and interpolation only drop slightly compared to the random split.
For Shapes3D (shown in the Appendix E), the OOD generalization is partially successful, especially in the composition and interpolation settings. We hypothesize that this is due to the dataset specific, fixed spatial composition of the images. For instance, with the object-centric positioning, the floor, wall and other factors are mostly at the same position within the images. Thus, they can reliably be inferred by only looking at a certain fixed spot in the image. In contrast, for MPI3D this is more difficult as, e.g., the robot finger has to be found to infer its tip color. Furthermore, the factors of variation in Shapes3D mostly consist of colors which are encoded within the same input dimensions, and not across pixels as, for instance, x-translation in MPI3D. For this color interpolation, the ReLU activation function might be a good inductive bias for generalization. However, it is not sufficient to achieve extrapolations, as we still observe a large drop in performance here.
Conclusion: The performance generally decreases when factors are OOD regardless of the supervision signal and architecture. However, we also observed exceptions in Shapes3D where OOD generalization was largely successful except for extrapolation.
2 Errors stem from inferring OOD factors
While in the previous section we observed a general decrease in score for the interpolation and extrapolation splits, our evaluation does not yet show how errors are distributed among individual factors that are in- and out-of-distribution.
In contrast to the previous section where multiple factors could be OOD distribution simultaneously, here, we control data splits (Figure 2) interpolation, extrapolation s.t. only a single factor is OOD. Now, we also estimate the -score separately per factor, depending on whether they have individually been observed during training (ID factor) or are exclusively in the test set (OOD factor). For instance, if we only have images of a heart with varying scale and position, we query the model with hearts at larger scales than observed during training (OOD factor), but at a previously observed position (ID factor). For a formal description, see Appendix Section H.2. This controlled setup enables us to investigate the modularity of the tested models, as we can separately measure the performance on OOD and ID factors. As a reference for an approximate upper bound, we additionally report the performance of the model on a random train/test split.
In Figures 5 and 14, we observe significant drops in performance for the OOD factors compared to a random test-train split. In contrast, for the ID factors, we see that the models still perform close to the random split, although with much larger variance. For the interpolation setting (Appendix Figure 14), this drop is also observed for MPI3D and dSprites but not for Shapes3D. Here, OOD and ID are almost on par with the random split. Note that our notion of modularity is based on systematic splits of individual factors and the resulting outputs. Other works focus on the inner behavior of a model by, e.g., investigating the clustering of neurons within the network (Filan et al. 2021). Preliminary experiments showed no correlations between the different notions of modularity.
Conclusion: The tested models can be fairly modular, in the sense that the predictions of ID factors remain accurate. The low OOD performances mainly stem from incorrectly extrapolated or interpolated factors. Given the low inter-/extrapolation (i.e., OOD) performances on MPI3D and dSprites, evidently no model learned to invert the ground-truth generative mechanism.
3 Models extrapolate similarly and towards the mean
In the previous sections, we observed that our tested models specifically extrapolate poorly on OOD factors. Here, we focus on quantifying the behavior of how different models extrapolate.
To check whether different models make similar errors, we compare the extrapolation behavior across architectures and seeds by measuring the similarity of model predictions for the OOD factors described in the previous section. No model is compared to itself if it has the same random seed. On MPI3D, Shapes3D and dSprites, all models strongly correlate with each other (Pearson ) but anti-correlate compared to the ground-truth prediction (Pearson ), the overall similarity matrix is shown in Appendix Fig. 17. One notable exception is on CelebGlow. Here, some models show low but positive correlations with the ground truth generative model (Pearson ). However, visually the models are still quite off as shown for the model with the highest correlation in Figure 18. In most cases, the highest similarity is along the diagonal, which demonstrates the influence of the architectural bias. This result hints at all models making similar mistakes extrapolating a factor of variation.
We find that models collectively tend towards predicting the mean for each factor in the training distribution when extrapolating. To show this, we estimate the following ratio of distances
where is the mean of FoV . If values of (2) are , models predict values which are closer to the mean than the corresponding ground-truth. We show a histogram over all supervised and transfer-based models for each dataset in Fig. 6. Models tend towards predicting the mean as only few values are . This is shown qualitatively in Appendix Figures 15 and 16.
Conclusion: Overall, we observe only small differences in how the tested models extrapolate, but a strong difference compared to the ground-truth. Instead of extrapolating, all models regress the OOD factor towards the mean in the training set. We hope that this observation can be considered to develop more diverse future models.
4 On the relation between disentanglement and downstream performance
Previous works have focused on the connection between disentanglement and OOD downstream performance (Träuble et al. 2020; Dittadi et al. 2020; Montero et al. 2021). Similarly, for our systematic splits, we measure the degree of disentanglement using the DCI-Disentanglement (Eastwood & Williams 2018) score on the latent representation of the embedded test and train data. Subsequently, we correlate it with the -performance of a supervised readout model which we report in Section 5.1. Note that the simplicity of the readout function depends on the degree of disentanglement, e.g., for a perfect disentanglement up to permutation and sign flips this would just be an assignment problem. For the disentanglement models, we consider the un-/ weakly supervised models -VAE(Higgins et al. 2017), SlowVAE (Klindt et al. 2020), Ada-GVAE(Locatello et al. 2020a) and PCL (Hyvarinen & Morioka 2017).
We find that the degree of downstream performance correlates positively with the degree of disentanglement (Pearson , Spearman ). However, the correlations vary per dataset and split (see Appendix fig. 7). Moreover, the overall performance of the disentanglement models followed by a supervised readout on the OOD split is lower compared to the supervised models (see e.g. Figure 4). In an ablation study with an oracle embedding that disentangles the test data up to permutations and sign flips, we found perfect generalization capabilities ().
Conclusion: Disentanglement models show no improved performance in OOD generalization. Nevertheless, we observe a mostly positive correlation between the degree of disentanglement and the downstream performance.
Other related benchmark studies
In this section, we focus on related benchmarks and their conclusions. For related work in the context of inductive biases, we refer to Section 3.
Corruption benchmarks: Other current benchmarks focus on the performance of models when adding common corruptions (denoted by -C) such as noise or snow to current dataset test sets, resulting in ImageNet-C, CIFAR-10-C, Pascal-C, Coco-C, Cityscapes-C and MNIST-C (Hendrycks & Dietterich 2019; Michaelis et al. 2019; Mu & Gilmer 2019). In contrast, in our benchmark, we assure that the factors of variations are present in the training set and merely have to be generalized correctly. In addition, our focus lies on identifying the ground truth generative process and its underlying factors. Depending on the task, the requirements for a model are very different. E.g., the ImageNet-C classification benchmark requires spatial invariance, whereas regressing factors such as, e.g., shift and shape of an object, requires in- and equivariance.
Abstract reasoning: Model performances on OOD generalizations are also intensively studied from the perspective of abstract reasoning, visual and relational reasoning tasks (Barrett et al. 2018; Wu et al. 2019; Santoro et al. 2017; Villalobos et al. 2020; Zhang et al. 2016; Yan & Zhou 2017; Funke et al. 2021; Zhang et al. 2018). Most related, (Barrett et al. 2018; Wu et al. 2019) also study similar interpolation and extrapolation regimes. Despite using notably different tasks such as abstract or spatial reasoning, they arrive at similar conclusions: They also observe drops in performance in the generalization regime and that interpolation is, in general, easier than extrapolation, and also hint at the modularity of models using distractor symbols (Barrett et al. 2018). Lastly, posing the concept of using correct generalization as a necessary condition to check whether an underlying mechanism has been learned has also been proposed in (Wu et al. 2019; Zhang et al. 2018; Funke et al. 2021).
Disentangled representation learning: Close to our work, Montero et al. (Montero et al. 2021) also study generalization in the context of extrapolation, interpolation and a weak form of composition on dSprites and Shapes3D, but not the more difficult MPI3D-Dataset. They focus on reconstructions of unsupervised disentanglement algorithms and thus the decoder, a task known to be theoretically impossible(Locatello et al. 2018). In their setup, they show that OOD generalization is limited. From their work, it remains unclear whether the generalization along known factors is a general problem in visual representation learning, and how neural networks fail to generalize. We try to fill these gaps. Moreover, we focus on representation learning approaches and thus on the encoder and consider a broader variety of models, including theoretically identifiable approaches (Ada-GAVE, SlowVAE, PCL), and provide a thorough in-depth analysis of how networks generalize.
Previously, Träuble et al. 2020 studied the behavior of unsupervised disentanglement models on correlated training data. They find that despite disentanglement objectives, the learned latent spaces mirror this correlation structure. In line with our work, the results of their supervised post-hoc regression models on Shapes3D suggest similar generalization performances as we see in our respective disentanglement models in Figures 4 and 11. OOD generalization w.r.t. extrapolation of one single FoV is also analyzed in (Dittadi et al. 2020). Our experimental setup in section 5.4 is similar to their ‘OOD2’ scenario. Here, our results are in accordance, as we both find that the degree of disentanglement is lightly correlated with the downstream performance.
Others: To demonstrate shortcuts in neural networks, Eulig et al. 2021 introduce a benchmark with factors of variations such as color on MNIST that correlate with a specified task but control for those correlations during test-time. In the context of reinforcement learning, Packer et al. 2018 assess models on systematic test-train splits similar to our inter-/extrapolation and show that current models cannot solve this problem. For generative adversarial networks (GANs), it has also been shown that their learned representations do not extrapolate beyond the training data (Jahanian et al. 2019).
Discussion and conclusion
In this paper, we highlight the importance of learning the independent underlying mechanisms behind the factors of variation present in the data to achieve generalization. However, we empirically show that among a large variety of models, no tested model succeeds in generalizing to all our proposed OOD settings (extrapolation, interpolation, composition). We conclude that the models are limited in learning the underlying mechanism behind the data and rather rely on strategies that do not generalize well. We further observe that while one factor is out-of-distribution, most other in-distribution factors are inferred correctly. In this sense, the tested models are surprisingly modular.
To further foster research on this intuitively simple, yet unsolved problem, we release our code as a benchmark. This benchmark, which allows various supervision types and systematic controls, should promote more principled approaches and can be seen as a more tractable intermediate milestone towards solving more general OOD benchmarks. In the future, a theoretical treatment identifying further inductive biases of the model and the necessary requirements of the data to solve our proposed benchmark should be further investigated.
The authors thank Steffen Schneider, Matthias Tangemann and Thomas Brox for their valuable feedback and fruitful discussions. The authors would also like to thank David Klindt, Judy Borowski, Dylan Paiton, Milton Montero and Sudhanshu Mittal for their constructive criticism of the manuscript. The authors thank the International Max Planck Research School for Intelligent Systems (IMPRS-IS) for supporting FT and LS. We acknowledge support from the German Federal Ministry of Education and Research (BMBF) through the Competence Center for Machine Learning (TUE.AI, FKZ 01IS18039A) and the Bernstein Computational Neuroscience Program Tübingen (FKZ: 01GQ1002). WB acknowledges support via his Emmy Noether Research Group funded by the German Science Foundation (DFG) under grant no. BR 6382/1-1 as well as support by Open Philantropy and the Good Ventures Foundation. MB and WB acknowledge funding from the MICrONS program of the Advanced Research Projects Activity (IARPA) via Department of Interior/Interior Business Center (DoI/IBC) contract number D16PC00003.
References
Ethics Statement
Our current study focuses on basic research and has no direct application or societal impact. Nevertheless, we think that the broader topic of generalization should be treated with great care. Especially oversimplified generalization and automation without a human in the loop could have drastic consequences in safety critical environments or court rulings.
Large-scale studies require a lot of compute due to multiple random seeds and exponentially growing sets of possible hyperparameter combinations. Following claims by Strubel et al. (Strubell et al. 2019), we tried to avoid redundant computations by orienting ourselves on current common values in the literature and by relying on systematic test runs. In a naive attempt, we tried in to estimate the power consumption and greenhouse gas impact based on the used cloud compute instance. However, too many factors such as external thermal conditions, actual workload, type of power used and others are involved (Mytton 2020; Fahad et al. 2019). In the future, especially with the trend towards larger network architectures, compute clusters should be required to enable options which report the estimated environmental impact. However, it should be noted that cloud vendors are already among the largest purchasers of renewable electricity (Mytton 2020).
For an impact statement for the broader field of representation learning, we refer to Klindt et al. 2020.
Reproducibility Statement
All important details to reproduce our results are repeated in Appendix H.
Appendix A Connection between readout performance and disentanglement of the representation
Here, we narrow down the root cause of the limited extrapolation performance of disentanglement models in the OOD settings as observed in Figures 4 and 11. More precisely, we investigate how the readout-MLP would perform on a perfectly disentangled representation. Therefore, we train our readout MLP directly on the ground-truth factors of variation for all possible test-train splits described in Figure 2 and measured the -score test error for each split. Here, the MLP only has to learn the identity function. In a slightly more evolved setting, termed sign-flip, we switched the sign input to train the readout-MLP on a mapping from -ground-truth to ground-truth. This mimics the identifiability guarantees of models like SlowVAE which are up to permutation and sign flips under certain assumptions. The R-squared for all settings in Table 1 are , therefore the readout model should not be the limitation for OOD generalization in our setting if the representation is identified up to permutation and sign flips. Note that this experiment does not cover disentanglement up to point-wise nonlinearities or linear/ affine transformations as required by other models.
Appendix B CelebGlow Dataset
The current disentanglement datasets such as dSprites, Shapes3D, MPI3D, and others are constructed based on highly controlled environments (Matthey et al. 2017; Kim & Mnih 2018; Gondal et al. 2019). Here, common factors of variations are rotations or scaling of simple geometric objects, such as a square. For a more intuitive investigation of other factors, we created the CelebGlow dataset. Here, the factors of variations are smiling, blondness and age. Samples are shown in Figure 8. Note that we rely on the Glow model instead of taking a real-world dataset, as this allows for a gradual control of individual factors of variation.
The CelebGlow dataset is created based on the invertible generative neural network of Kingma et al. (Kingma & Dhariwal 2018). We used their provided network The network can be found at: https://github.com/openai/glow/blob/master/demo/script.sh#L24 that is pretrained on the Celeb-HQ dataset, and has labelled directions in the model-latent space that correspond to specific attributes of the dataset. Based on this latent space, we created the dataset as follows:
In the latent space of the model, we sample from a high dimensional Gaussian with zero mean and a low standard deviation of 0.3 to avoid too much variability.
Next, we perform a latent walk into the directions that correspond to "Smiling", "Age" and "Blondness" in image space. To estimate the spacing, we rely on the function manipulate_range https://github.com/openai/glow/blob/master/demo/model.py#L219. We perform 6 steps along each axis and all combinations (6x6x6 cube). As a scale parameter to the function, we use 0.8. Those factors were chosen s.t. the images differ significantly, but also to stay in the valid range of the model based on visual inspection.
We pass all latent coordinates through the glow network in the generative direction.
We further down-sample the images from 256x256x3 to 64x64x3 to match the resolution of common disentanglement datasets.
Finally, we store each image and the corresponding factor combination.
This procedure is repeated for 1000 samples to get samples in total, which is around the same size as other common datasets.
Appendix C Hyperparameter Tuning Ablation
As described in the implementation details, we use common values from the literature to train the proposed models. Here, we investigate effects of such hyperparameters on the CNN architecture. Due to the combinatorial complexity, we do not perform a search for other architectures. As hyperparameters, we varied the number or training iterations (3 different numbers of iterations), we introduced 5 different strengths of regularization, 2 different depths for the CNN architecture [6 layers, 9 layers] and ran multiple random seeds for each combination.
The results on the extrapolation test on MPI3D set are shown in Figure 9. Given this hyperparameter search, we find no improvement over our reported numbers for the CNN.
Appendix D Real versus synthetic dataset
To narrow down the question “why the generalization capabilities drop on real-world dataset MPI3D?”, we run a comparison on MPI3D dataset with real and synthetic images.
The results on the MPI3D dataset with synthetic images is shown in Figure 10 and Table 10. Comparing this with R-squared performances to MPI3D with real-world images (Figure 4 and Table 9), we observe that the results do not change significantly (most results are in a 1-2sigma range). We conclude that the larger drops in performance on MPI3D compared to Shapes3D or dSprites, are not due to the real images as opposed to synthetic images. Instead, we hypothesize that it is due to the more realistic setup of the MPI3D dataset itself. For instance, it contains complex factors like rotation in 3D projected on 2D. Here, occluded parts of objects have to be guessed based on certain symmetry assumptions.
Appendix E Ablation on non-ambiguous dSprites
The setup of dSprites is non-injective, as different rotations map to the same image. E.g., the square at a rotation is identical to the one rotated by and therefore ambiguous. Thus, the training process is noisy. In an ablation study, we controlled for this by constraining the rotations to lie in . We again ran all our proposed models and report the -Score in Figures 11(b) and 11(b).
Comparing the new results with the original dSprites results shows: First, for the random test-train split, resolving the rotational ambiguity leads to almost perfect performance (close to R-squared scores for most models). In the previous dSprites setup with rotational ambiguity, top accuracies are around 70-95% R-squared scores for most models. Second, large drops in performance can still be observed when we move towards the systematic out-of-distribution splits (composition, interpolation, and extrapolation). Also, our insights on how models extrapolate remain the same. Lastly, for the random split, the Rotation-EQ model shows non-perfect performance. Tracing this error to individual factors, it turns out this is due to limited capabilities in predicting the x, y positions. We hypothesize that this is due to limitations of convolutions in propagating spatial positions, as discussed in (Liu et al. 2018). The DenseNet performs perfectly on the train set and might be overfitting.
We conclude that the rotational ambiguity explains the drops on the random split. However, the clear drops in performance on the systematic splits remain nonetheless. Thus, the analysis we perform in the paper and the conclusions we draw remain the same.
Appendix F Data Augmentations
We investigate the effects of data augmentation during training time on the generalization performance in the extrapolation setting of our proposed benchmark.
As data augmentations, we applied random erasing, Gaussian Noise, small shearings, and blurring. Note that we could not use arbitrary augmentations. For instance, shift augmentations would lead to ambiguities with the “shift” factor in dSprites. Next, we trained CNNs with and without data augmentations on all four datasets (dSprites, Shapes3D, MPI3D, CelebGlow) on the extrapolation splits with multiple random seeds.
The results are visualized in Figure 12. For the mean performance, we observe no significant improvement by adding augmentations. However, the overall spread of the scores seems to decrease given augmentations on some datasets. We explain this by the fact that the augmentations enforce certain invariances, narrowing the solution space of optimal training solutions by providing a further specification (specification in the sense of D’Amour et al. 2020).
Appendix G Performance with respect to individual factors
We here try to attribute the performance losses to individual OOD factors (see Section 5.2). Thus, on the extrapolation setting, we modify the test-splits such that only a single factor is out-of-distribution. Next, we measure the overall performance across models (all fully supervised and transfer models) to demonstrate the effect of this factor. The results are depicted for all models in Figure 13. Overall, factors like "height" on MPI3D that control the viewing of the camera and, subsequently, change attributes like the absolute position in the image of other factors (e.g., the tip of the robot arm) have a high effect.
Appendix H Implementation details
Each dataset consists of multiple factors of variation and every possible combination of factors generates a corresponding image. Here, we list all datasets and their corresponding factor ranges. Note, to estimate the reported -score, we normalize the factors by dividing each factor by , i.e., all factors are in the range $$. dSprites (Matthey et al. 2017), represents some low resolution binary images of basic shapes with the 5 FoVs shape {0, 1, 2}, scale {0, …, 4}, orientation Note that this dataset contains a non-injective generative model as square and ellipses have multiple rotational symmetries. {0, …, 39}, x-position {0, …, 31}, and y-position {0, …, 31}. Next, Shapes3D (Kim & Mnih 2018) which is a similarly popular dataset with 3D shapes in a room scenes defined by the 6 FoVs floor color {0, …, 9}, wall color {0, …, 9}, object color {0, …, 9}, object size {0, …, 7}, object type {0, …, 3} and azimuth {0, …, 14}. Lastly, we consider the challenging and more realistic dataset MPI3D (Gondal et al. 2019) containing real images of physical 3D objects attached to a robotic finger generated by 7 FoVs color {0, …, 5}, shape {0, …, 5}, size {0, 1}, height {0, 1, 2}, background color {0, 1, 2}, x-axis {0, …, 39} and y-axis {0, …, 39}.
H.2 Data set Splits
Each dataset is complete in the sense that it contains all possible combinations of factors of variation. Thus, the interpolation and extrapolation test-train splits are fully defined by specifying which factors are exclusively in the test set. Starting from all possible combinations, if a given factor value is defined to be exclusively in the test set, the corresponding image is part of the test set. E.g. for the extrapolation case in dSprites, all images containing x-positions > 24 are part of the test set and the train set its respective complement . Composition can be defined equivalently to extrapolation but with interchanged test and train sets. The details of the splits are provided in table Tables 2 and 3. The resulting train vs. test sample number ratios are roughly . See Table 4. We will release the test and train splits to allow for a fair comparison and benchmarking for future work.
For the setting where only a single factor is OOD, we formally define this as
H.3 Training
All models are implemented using PyTorch 1.7. If not specified otherwise, the hyperparameters correspond to the default library values.
For the un-/weakly supervised models, we consider 10 random seeds per hyperparameter setup. As hyperparameters, we optimize one parameter of the learning objective per model similar to Table 2 from Locatello et al. (Locatello et al. 2020a). For the SlowVAE, we took the optimal values from Klindt et al. (Klindt et al. 2020) and tuned for . The PCL model itself does not have any hyperparameters (Hyvarinen & Morioka 2017). For simplicity, we determine the optimal setup in a supervised manner by measuring the DCI-Disentanglement score (Eastwood & Williams 2018) on the training split. The PCL and SlowVAE models are trained on pairs of images that only differ sparsely in their underlying factors of variation following a Laplace transition distribution, the details correspond to the implementation https://github.com/bethgelab/slow_disentanglement/blob/master/scripts/dataset.py#L94 of Klindt et al. (Klindt et al. 2020). The Ada-GVAE models are trained on pairs of images that differ uniformly in a single, randomly selected factor. Other factors are kept fixed. This matches the strongest model from Locatello et al. (Locatello et al. 2020a) implemented on GitHub https://github.com/google-research/disentanglement_lib/blob/master/disentanglement_lib/methods/weak/weak_vae.py#L62 and https://github.com/google-research/disentanglement_lib/blob/master/disentanglement_lib/methods/weak/weak_vae.py#L317. All -VAE models are trained in an unsupervised manner. All un- and weakly supervised models are trained with the Adam optimizer with a learning rate of . We train each model for iterations with a batch size of 64, which for the weakly supervised models, corresponds to 64 pairs. Lastly, we train a supervised readout model on top of the latents for 8 epochs with the Adam optimizer on the full corresponding training dataset and observe convergence on the training and test datasets - no overfitting was observed.
Fully supervised:
Transfer learning:
The pre-trained models are fine-tuned with the same loss as the fully supervised models. We train for iterations and with a lower learning rate of . We fine-tune all model weights. As an ablation, we also tried only training the last layer while freezing the other weights. In this setting, we consistently observed worse results and, therefore, do not include them in this paper.
H.4 Model implementations
Here, we shortly describe the implementation details required to reproduce our model implementation. We denote code from Python libraries in grey. If not specified otherwise, the default parameters and nomenclature correspond to the PyTorch 1.7 library.
The un- and weakly supervised models -VAE, Ada-GVAE and SlowVAE all use the same encoder-decoder architecture as Locatello et al. (Locatello et al. 2020a). The PCL model uses the same architecture as the encoder as well and with the same readout structure for the contrastive loss as used by Hyvärinen et al. (Hyvarinen & Morioka 2017). For the supervised readout MLP, we use the sequential model [Linear(10, 40), ReLU(), Linear(40, 40), ReLU(40, 40), Linear(40, 40), ReLU(), Linear(40, number-factors)].
The MLP model consists of [Linear(64*64*number-channels, 90), ReLU(), Linear(90, 90), ReLU(), Linear(90, 90), ReLU(), Linear(90, 90), ReLU(), Linear(90, 45), ReLU(), Linear(22, number-factors)]. The architecture is chosen such that it has roughly the same number of parameters and layers as the CNN.
The CNN architecture corresponds the one used by Locatello et al. (Locatello et al. 2020a). We only adjust the number of outputs to match the corresponding datasets.
The CoordConv consists of a CoordConv2D layer following the PyTorch implementation https://github.com/walsvid/CoordConv with 16 output channels. It is followed by 5 ReLU-Conv layers with 16 in- and output channels each and a MaxPool2D layer. The final readout consists of [Linear(32, 32), ReLU(), Linear(32, number-factors)].
The SetEncoder concatenates each input pixel with its pixel coordinates normalized to $$. All concatenated pixels (i, j, pixel-value) are subsequently processed with the same network which consists of [Linear(2+number-channels), ReLU(), Linear(40, 40), ReLU(), Linear(40, 20), ReLU()]. This is followed by a mean pooling operation per image which guarantees an invariance over the order of the inputs, i.e. one could shuffle all inputs and the output would remain the same. As a readout, it follows a sequential fully connected network consisting of [Linear(20, 20), ReLU(), Linear(20, 20), ReLU(), Linear(20, number-factors)].
The rotationally equivariant network RotEQ is similar to the architecture from Locatello et al. (Locatello et al. 2020a). One difference is that it uses the R2Conv module https://github.com/QUVA-Lab/e2cnn from Weiler et al. (Weiler & Cesa 2019) instead of the PyTorch Conv2d with an 8-fold rotational symmetry. We thus decrease the number of feature maps by a factor of 8, which roughly corresponds to the same computational complexity as the CNN. We provide a second version which does not decrease the number of feature maps and, thus, has the same number of trainable parameters as the CNN but a higher computational complexity. We refer to this version as RotEQ-big.
To implement the spatial transformer (STN) (Wu et al. 2019), we follow the PyTorch tutorial implementation https://pytorch.org/tutorials/intermediate/spatial_transformer_tutorial.html which consists of two steps. In the first step, we estimate the parameters of a (2, 3)-shaped affine matrix using a sequential neural network with the following architecture [Conv2d(number_channels, 8, kernel_size=7), MaxPool2d(2, stride=2), ReLU(), Conv2d(8, 10, kernel_size=5), MaxPool2d(2, stride=2), ReLU(), Conv2d(10, 10, kernel_size=6), MaxPool2d(2, stride=2), ReLU(), Linear(10*3*3, 31), ReLU(), Linear(32, 3*2)]. In the second step, the input image is transformed by the estimated affine matrix and subsequently processed by a CNN which has the same architecture as the CNN described above.
For the transfer learning models ResNet50 (RN50) and ResNet101 (RN101) pretrained on ImageNet-21k (IN-21k), we use the big-transfer (Kolesnikov et al. 2020) implementation https://colab.research.google.com/github/google-research/big_transfer/blob/master/colabs/big_transfer_pytorch.ipynb and for the weights https://storage.googleapis.com/bit_models/{bit_variant}.npz. For the RN50, we download the weights with the tag "BiT-M-R50x1", and for the RN101, we use the tag "BiT-M-R101x3". For the DenseNet trained on ImageNet-1k (IN-1k), we used the weights from densenet121. For all transfer learning methods, we replace the last layer of the pre-trained models with a randomly initialized linear layer which matches the number of outputs to the number of factors in each dataset. As an ablation, we also provide a randomly initialized version for each transfer learning model.
H.5 Compute
All models are run on the NVIDIA T4 Tensor Core GPUs on the AWS g4dn.4xlarge instances with an approximate total compute of 20 000 GPUh. To save computational cost, we gradually increased the number of seeds until we achieved acceptable p-values of . In the end, we have 3 random seeds per supervised model and 10 random seeds per hyperparameter setting for the un and weakly supervised models.