Robustly Disentangled Causal Mechanisms: Validating Deep Representations for Interventional Robustness
Raphael Suter, Đorđe Miladinović, Bernhard Schölkopf, Stefan Bauer
Introduction
Learning deep representations in which different semantic aspects of data are structurally disentangled is of central importance for training robust machine learning models. Separating independent factors of variation could pave the way for successful transfer learning and domain adaptation (Bengio et al., 2013). Imagine the example of a robot learning multiple tasks by interacting with its environment. For data efficiency, the robot can learn a generic representation architecture that maps its high dimensional sensory data to a collection of general, compact features describing its surrounding. For each task, only a subset of features will be required. If the robot is instructed to grasp an object, it must know the shape and the position of the object, however, its color is irrelevant. On the other hand, when pointing to all red objects is demanded, only position and color are required.
Having a disentangled representation, where each feature captures only one factor of variation, allows the robot to build separate (simple) models for each task based on only a relevant and stable subselection of these generically learned features. We argue that robustness of the learned representation is a crucial property when this is attempted in practice. It has been proposed that features should be selected based on their robustness or invariance across tasks (e.g., Rojas-Carulla et al., 2018), we hence do not want them to be affected by changes in any other factor. In our example, the robot assigned with the grasping task should be able to build a model using features well describing shape and position of the object. For this model to be robust, however, these features must not be affected by changing color (or any other nuisance factor).
It is striking that despite the recent popularity of disentangled representation learning approaches, a commonly accepted definition and validation metric is missing (Higgins et al., 2018). We view disentanglement as a property of a causal process (Spirtes et al., 1993; Pearl, 2009) responsible for the data generation, as opposed to only a heuristic characteristic of the encoding. Concretely, we call a causal process disentangled when the parents of the generated observations do not affect each other (i.e., there is no total causal effect between them (Peters et al., 2017, Definition 6.12)). We call these parents elementary ingredients. In the example above, we view color and shape as elementary ingredients, as both can be changed without affecting each other. Still, there can be dependencies between them if for example our experimental setup is confounded by the capabilities of the 3D printers that are used to create the objects (e.g., certain shapes can only be printed in some colors).
Combining these disentangled causal processes with the encoding allows us to study interventional effects on feature representations and estimate them from observational data. This is of interest when benchmarking disentanglement approaches based on ground truth data (Locatello et al., 2018) or trying to evaluate robustness of a deep representations w.r.t. known nuisance factors (e.g., domain changes). In the example of robotics, knowledge about the generative factors (e.g., the color, shape, weight, etc. of an object to grasp) is often availabe and can be controlled in experiments.
We will start by first giving an overview of previous work in finding disentangled representations and how they have been validated in Section 2. In Section 3 we introduce our framework for the joint treatment of the disentangled causal process and its learned representation. We introduce our notion of interventional effects on encodings and the following interventional robustness score in Section 4 and show how this score can be estimated from observational data with an efficient algorithm in Section 5. Section 6 provides experimental evidence in a standard disentanglement benchmark dataset supporting the need of a robustness based disentanglement criterion.
We introduce a unifying causal framework of disentangled generative processes and consequent feature encodings. This perspective allows us to introduce a novel validation metric, the interventional robustness score.
We show how this metric can be estimated from observational data and provide an efficient algorithm that scales linearly in the dataset size.
Our extensive experiments on a standard benchmark dataset show that our robustness based validation is able to discover vulnerabilities of deep representations that have been undetected by existing work.
Motivated by this metric, we additionally present a new visualisation technique which provides an intuitive understanding of dependency structures and robustness of learned encodings.
Notation:
We denote the generative factors of high dimensional observations as . The latent variables learned by a model, e.g., a variational auto-encoder (VAE) (Kingma & Welling, 2014), are denoted as . We use the notation to describe the encoding which in case of VAEs corresponds to the posterior mean of . Capital letters denote random variables, and lower case observations thereof. Subindices for a set or for a single index denote the selected components of a multidimensional variable. A backslash denotes all components except those in .
Related Work
In the framework of variational auto-encoders (VAEs) (Kingma & Welling, 2014) the (high dimensional) observations are modelled to be generated from some latent features with chosen prior according to the probabilistic model . The generative model as well as the proxy posterior can be estimated using neural networks by maximizing the variational lower bound (ELBO) of :
This objective function a priori does not encourage much structure on the latent space (except some similarity to the chosen prior which is usually isotropic Gaussian). More precisely, for a given encoder and decoder any bijective transformation of the latent space yields the same reconstruction .
Various proposals for more structure imposing regularization have been made, either with some sort of supervision (e.g. Siddharth et al., 2017; Bouchacourt et al., 2017; Liu et al., 2017; Mathieu et al., 2016; Cheung et al., 2014) or completely unsupervised (e.g. Higgins et al., 2017; Kim & Mnih, 2018; Chen et al., 2018; Kumar et al., 2018; Esmaeili et al., 2018). Higgins et al. (2017) proposed the -VAE penalizing the Kullback-Leibler divergence (KL) term in the VAE objective (1) more strongly, which encourages similarity to the factorized prior distribution. Others used techniques to encourage statistical independence between the different components in , e.g., FactorVAE (Kim & Mnih, 2018) or -TCVAE (Chen et al., 2018), similar to independent component analysis (e.g. Comon, 1994). With disentangling the inferred prior (DIP-VAE), Kumar et al. (2018) proposed encouraging factorization of .
A special form of structure in the latent space which has gained a lot of attention in recent time is referred to as disentanglement (Bengio et al., 2013). This term encompasses the understanding that each learned feature in should represent structurally different aspects of the observed phenomena (i.e., capture different sources of variation).
Various methods to validate a learned representation for disentanglement based on known ground truth generative factors have been proposed (e.g. Eastwood & Williams, 2018; Ridgeway & Mozer, 2018; Chen et al., 2018; Kim & Mnih, 2018). While a universal definition of disentanglement is missing, the most widely accepted notion is that one feature should capture information of only one generative factor (Eastwood & Williams, 2018; Ridgeway & Mozer, 2018). This has for example been expressed as the mutual information of a single latent dimension with generative factors (Ridgeway & Mozer, 2018), where in the ideal case each has some mutual information with one generative factor but none with all the others. Similarly, Eastwood & Williams (2018) trained predictors (e.g., Lasso or random forests) for a generative factor based on the representation . In a disentangled model, each dimension is only useful (i.e., has high feature importance) to predict one of those factors (see appendix D for details).
Validation without known generative factors is still an open research question and so far it is not possible to quantitatively validate disentanglement in an unsupervised way. The community has been using ”latent traversals” (i.e., changing one latent dimension and subsequently re-generating the image) for visual inspection when supervision is not available (see e.g. Chen et al., 2018). This can be used to encounter physically meaningful interpretations of each dimension.
Causal Model
We will first consider assumptions for the causal process underlying the data generating mechanism. Following this, we discuss consequences for trying to match encodings with causal factors in a deep latent variable model.
As opposed to previous approaches that defined disentanglement heuristically as properties of the learned latent space, we take a step back and first introduce a notion of disentanglement on the level of the true causal mechanism (or data generation process). Subsequently, we can use this definition to better understand a learned probabilistic model for latent representations and evaluate its properties.
We assume to be given a set of observations from a (potentially high dimensional) random variable . In our model, the data generating process is described by causes of variation (generative factors) (i.e., ) that do not cause each other. These factors are generally assumed to be unobserved and are objects of interest when doing deep representation learning. In particular, knowledge about could be used to build lower dimensional predictive models, not relying on the (unstructured) itself. This could be classic prediction of a label , often in ”confounded” direction (i.e., predicting effects from other effects) if or in anti-causal direction if . It is also relevant in a domain change setting when we know that the domain has an impact on , i.e., .
Having these potential use cases in mind, we assume the generative factors themselves to be confounded by (multi-dimensional) , which can for example include a potential label or source . Hence, the resulting causal model allows for statistical dependencies between latent variables and , , when they are both affected by a certain label, i.e., .
However, a crucial assumption of our model is that these latent factors should represent elementary ingredients to the causal mechanism generating (to be defined below), which can be thought of as descriptive features of that can be changed without affecting each other (i.e., there is no causal effect between them). A similar assumption on the underlying model is likewise a key requirement for the recent extension of identificability results of non-linear ICA (Hyvarinen et al., 2018). We formulate this assumption of a disentangled causal model as follows (see also Figure 1):
In practice we assume that the dimensionality of the confounding is significantly smaller than the number of factors .
This definition reflects our understanding of elementary ingredients , of the causal process. Each ingredient should work on its own and is changable without affecting others. This reflects the independent mechanisms (IM) assumption (Schölkopf et al., 2012). Independent mechanisms as components of causal models allow intervention on one mechanism without affecting the other modules and thus correspond to the notion of independently controllable factors in reinforcement learning (Thomas et al., 2017). Our setting is broader, describing any causal process and inheriting the generality of the notion of IM, pertaining to autonomy, invariance and modularity (Peters et al., 2017).
Based on this view of the data generation process, we can prove (see Appendix B) the following observations which will help us discuss notions of disentanglement and deep latent variable models.
A disentangled causal process as introduced in Definition 1 fulfills the following properties:
describes a causal mechanism invariant to changes in the distributions .
In general, the latent causes can be dependent
Only if we condition on the confounders in the data generation they are independent
Knowing what observation of we obtained renders the different latent causes dependent, i.e.,
The latent factors already contain all information about confounders that is relevant for , i.e.,
where denotes the mutual information.
There is no total causal effect from to for ; i.e., intervening on does not change , i.e,
The remaining components of , i.e., , are a valid adjustment set (Pearl, 2009) to estimate interventional effects from to based on observational data, i.e.,
If there is no confounding, conditioning is sufficient to obtain the post interventional distribution of :
2 Disentangled Latent Variable Model
We can now understand generative models with latent variables (e.g., the decoder in VAEs) as models for the causal mechanism in a and the inferred latent space through as proxy to the generative factors . Property d gives hope that under an adequate information bottleneck we can indeed recover information about causal parents and not the confounders. Ideally, we would hope for a one-to-one correspondance of to for all . In some situations it might be useful to learn multiple latent dimensions for one causal factor for a more natural description, e.g., describing an angle as and (Ridgeway & Mozer, 2018). Hence, we will generally allow the encodings to be dimensional, where usually . The -VAE (Higgins et al., 2017) encourages factorization of through penalization of the KL to its prior . Due to property c other approaches were introduced making use of statistical independence (Kim & Mnih, 2018; Chen et al., 2018; Kumar et al., 2018). Esmaeili et al. (2018) allow dependence within groups of variables in a hierarchical model (i.e., with some form of confounding where property b becomes an issue) by specifically modelling groups of dependent latent encodings. In contrast to the above mentioned approaches, this requires prior knowledge on the generative structure. We will make use of property f to solve the task of using observational data to evaluate deep latent variable models for disentanglement and robustness.
Figure 2 illustrates our causal perspective on representation learning which encompasses the data generating process () as well as the subsequent encoding through (). Based on this viewpoint, we define the interventional effect of a group of generative factors on the implied latent space encodings with proxy posterior from a VAE, where and as:
This definition is consistent with the above graphical model as it implies that .
Interventional Robustness
Building on the definition of interventional effects on deep feature representations in Eq. (3.2), we now derive a robustness measure of encodings with respect to changes in certain generative factors.
Let and be groups of indices in the latent space and generative space. For generality, we will henceforth talk about robustness of groups of features with respect to interventions on groups of generative factors . We believe that having this general formulation of allowing disagreements between groups of latent dimensions and generative factors provides more flexibility, for example when multiple latent dimensions are used to describe one phenomenon (Esmaeili et al., 2018) or when some sort of supervision is available through groupings in the dataset according to generative factors (Bouchacourt et al., 2017). Below, we will also discuss special cases of how these sets can be chosen.
If we assume that the encoding captures information about the causal factors and we would like to build a predictive model that only depends on those factors, we might be interested in knowing how robust our encoding is with respect to nuisance factors , where . To quantify this robustness for specific realizations of and we make the following definition:
For any given set of feature indices , and , we call
One important special case includes the setting where , and . This corresponds to the degree to which is robustly isolated from any extraneous causes (assuming captures ), which can be interpreted as the concept of disentanglement in the framework of Eastwood & Williams (2018). We define
as disentanglement score of . The maximizing is interpreted as the generative factor that captures predominantly. Intuitively, we have robust disentanglement when a feature reliably captures information about the generative factor , where reliable means that the inferred value is always the same when stays the same, regardless of what the other generative factors are doing.
quantifies how robust is when changes in occur. If we are building a model predicting a label based on some (to be selected) feature set , we can use this score to make a trade-off between robustness and predictive power. For example, we could use the best performing set of features among all those that satisfy a given robustness threshold.
Estimation and Benchmarking Disentanglement
The proof of Proposition 2 can be found in Appendix C. Note that a dataset capturing all possible variations generally grows exponentially in the number of generative factors. While this is a general issue for all validation approaches and care needs to be taken when collecting such datasets in practice, we just remind that due to the generally large nature of it is particularly important to have such an efficient validation procedure. In many benchmark datasets for disentanglement (e.g. dsprites) the observations are obtained noise-free and the dataset contains all possible combinations of generative factors exactly once. This makes the estimation of the disentanglement score even easier, as we have . Furthermore, since no confounding is present, we can use conditioning to estimate the interventional effect, i.e., , as seen in Proposition 1 g. The disentanglement score of , as discussed in Eq. (3) , follows (see A.1 for details) as:
Experiments
Our evaluations involve five different state of the art unsupervised disentanglement techniques (classic VAE, -VAE, DIP-VAE, FactorVAE and -TCVAE), each learning features.
Believing that it is most insightful to look at scores for each dimension separately, which indicates the quality of a single feature, we included the full evaluations including plots of correspondance matrices (as in Figure 6) in Appendix E. For future extensions and applications our work is added to the disentanglement_lib of Locatello et al. (2018).
2 Robustness as Complementary Metric
3 Visualising Interventional Robustness
Conclusion
We have proposed a framework for assessing disentanglement in deep representation learning which combines the generative process responsible for high dimensional observations with the subsequent feature encoding by a neural network. This perspective leads to a natural validation method, the interventional robustness score. We show how it can be estimated from observational data using an efficient algorithm that scales linearly in the dataset size. As special cases, this proposed measure captures robust disentanglement and domain shift stability. Extensive evaluations showed that the existing metrics do not capture the effects that rare events or cumulative influences from multiple generative factors can have on feature encodings, while our robustness based validation metric discovers such vulnerabilities.
We envision that the notion of interventional effects on encodings may give rise to the development of novel, robustly disentangled representation learning algorithms, for example in the interactive learning environment (Thomas et al., 2017) or when weak forms of supervision are available (Bouchacourt et al., 2017; Locatello et al., 2019). The exploration of those ideas, especially including confounding, is left for future research.
Acknowledgments
We thank Andreas Krause for helpful discussions and Alexander Neitz, Francesco Locatello and Olivier Bachem for the inclusion of our work into disentanglement_lib (https://github.com/google-research/disentanglement_lib) of Locatello et al. (2018). This research was partially supported by the Max Planck ETH Center for Learning Systems.
References
Appendix A Estimation
The main ingredient for this estimation to work is provided by our constrained causal model (i.e., a disentangled process) that implies that the backdoor criteria can be applied, which we showed in Proposition 1. Further, we already saw in Eq. (3.2) that . This can be used to write the conditional expected value of as:
where the elements of encoding are defined as:
By denoting the Kronecker-delta as we obtain:
which gives us the natural interpretation that samples that would occur more often together with a certain need to be downweighted in order to correct for the confounding effects. We can also see that in case of statistical independence between the generative factors, this reweighting is not needed and we can simply use the sample mean with the subselection of the dataset .
The estimate for the disentanglement score in Eq. (3) for follows from that:
Appendix B Proof of Proposition 1
Property a directly follows from Definition 1 and the definition of an independent causal mechanism. b and c can be read off the graphical model (Koller et al., 2009) in Figure 1 which does not contain any arrow from to for by Definition 1 of the constrained SCM. This is due to the fact that any distribution implied by an SCM is Markovian with respect to the corresponding graph (Peters et al., 2017, Prop. 6.31). d follows from the data processing inequality since we have . The non-existence of a directed path from to implies that there is no total causal effect (Peters et al., 2017, Prop. 6.14). This, in turn, is equivalent to property e (Peters et al., 2017, Prop. 6.13). Finally, since there are no arrows between the ’s, the backdoor criterion (Peters et al., 2017, Prop. 6.41) can be applied to estimate the interventional effects in f. In particular, blocks all paths from to entering through the backdoor (i.e., ) but at the same time does not contain any descendents of since by definition . Property g also follows from by using parent adjustment (Peters et al., 2017, Prop. 6.41), where in the case no confounding . These properties is why the constrained SCM in Definition 1 is important for further estimation. ∎
Appendix C Proof of Proposition 2
The encodings in line 6 requires one pass through the dataset . So does the estimation of the occurance frequencies in line 7 as one can use a hash table to keep track of the number of occurances of each possible realization. Therefore, the preprocessing steps scale with .
Further, also the partitioning of the full dataset into , which is done in lines 9, 10 and 13, can be done with two passes through the dataset by using hash tables: In the first pass we create buckets with as keys. Consequently, we can pass through all of these buckets to create subbuckets where is used as key. This reasoning is further illustrated in Figure 5 and leads us to the complexity of the partitioning.
Though this estimation procedure scales in the dataset size, the required number of observations for a fixed estimation quality (i.e., if should stay constant) might become very large, as we have exponentially growing (in and ) many possible combinations to consider. This is why some trade-offs need to be made when comparing large sets of factors. The estimation for , however, usually works well. One trade-off parameter is the discretization step of of ’s. Partitioning a factor into fewer realizations yields less possible combinations and hence larger sets . In general, the more noise we expect in the larger the sets we want to have in order to obtain stable estimates of the expected values. Also, if we allow for fewer possible realizations in the generative factors, the smaller our dataset can be to cover all relevant combinations. However, larger discretization steps come at the cost of having a less sensitive score. Also note that taking the supremum is in general not vulnerable to outliers in as we compute distances of expected values. When outliers are to be expected, a robust estimate for these expected values can be used. Only when little data is available special care needs to be taken.
Appendix D Details of Experimental Setup
We compute the feature importance based disentanglement scores, as discussed by Eastwood & Williams (2018), using random forests with 50 decision trees that are split up to a minimal leaf size of 500. As opposed to Eastwood & Williams (2018), we only use one single feature to ’randomly choose from’ at each split, since this guarantees that each feature is equally given the chance to prove itself in reducing the out-of-bag error. When multiple features can be chosen from at each split, it is well possible that features with a mediocre importance are never chosen as there are features always yielding a better split. This would lead to an underestimation of their importance.
For the mutual information metric (Ridgeway & Mozer, 2018) we followed the original proposal of discretizing each latent dimension into 20 buckets and computing the discrete mutual information based on that. We found that using smaller discretization steps (i.e., more buckets) does not change the results notably.
Plotting the matrix gives a good first impression of the disentanglement capabilities of an encoder. Ideally, we would want to see only one large value per row while the remaining entries should be zero. In our experimental evaluations we plot this matrix (together with similarly interpretable matrices of the other metrics) as is shown for example in Figure 6 on page 6.
To explicitly quantify this visual perspective, Eastwood & Williams (2018) summarize disentanglement as one score value which measures to what extent indeed each latent dimension can only be used to predict one generative factor (i.e., sparse rows). It is obtained by first computing the ‘probabilities’ of being important to predict ,
and the entropy of this distribution: , where is the number of generative factors. The disentanglement score of variable is then defined as For example, if only one generative factor can be predicted with , i.e., , we obtain . If the explanatory power spreads over all factors equally, the score is zero. Using relative variable importance , which accounts for dead or irrelevant components in , they find an overall disentanglement score as weighted average . When later plotting the full importance matrices, we also provide information about the individual feature disentanglement scores in the corresponding row labels. These feature-wise scores are better comparable between metrics since all of them have different heuristics to obtain the (weighted) average .
As an additional measure to obtain a more complete picture of the quality of the learned code, they additionally propose the informativeness score. It tells us how much information about the generative factors is captured in the latent space and is computed as the out-of-bag prediction accuracy of the regressors . In our evaluations in Section 6 we will also provide this score, as there is often a trade-off between a disentangled structure and information being preserved.
D.2 Disentanglement Approaches
For the disentangling VAE models we made use of existing implementations where this was available. Classic VAE (Kingma & Welling, 2014) and DIP-VAE (Kumar et al., 2018) we implemented ourselves and trained them for epochs using Adam (Kingma & Ba, 2015) with a learning rate of 1e-4 and batch size of . We used the same neural network architecture as is described in the appendix of Chen et al. (2018). For DIP-VAE we set the parameters to , as is used in the original publication. For the annealed -VAE approach (Burgess et al., 2018) we used the publicly available third party code from https://github.com/1Konny/Beta-VAE, where parameters are set to and . Also, for FactorVAE (Kim & Mnih, 2018) we used third party code from https://github.com/1Konny/FactorVAE with their parameter . Chen et al. (2018) provided their own code for -TCVAE at https://github.com/rtqichen/beta-tcvae, which we made use of. We kept their chosen default parameters ().
Appendix E Visualisations of Importance Matrices
Plots of the full importance matrices for the considered latent spaces and all three validation metrics are included in Figures 8, 9, 10, 11 and 12. The y labels include the disentanglement scores of each individual feature .
A related visualization possibility to the one we propose in Section 6.3 is that of simple conditioning on different generative factors (without keeping one factor fixed). This is illustrated in Figure 7, where we plot the violin plots (i.e., density estimates) of for all generative factors (columns) and realizations of them (x axis). This kind of visualization works well to discover simple dependency patterns as well as their noise levels.
Appendix F Visualisations of Interventional Effects
We provide further visualizations of the full latent spaces and their dependency structure (produced by the to be made publicly available code) of a couple of models in Figures 13, 14, 15, 16 and 17.