Variational Inference of Disentangled Latent Concepts from Unlabeled Observations

Abhishek Kumar, Prasanna Sattigeri, Avinash Balakrishnan

Introduction

Feature representations of the observed raw data play a crucial role in the success of machine learning algorithms. Effective representations should be able to capture the underlying (abstract or high-level) latent generative factors that are relevant for the end task while ignoring the inconsequential or nuisance factors. Disentangled feature representations have the property that the generative factors are revealed in disjoint subsets of the feature dimensions, such that a change in a single generative factor causes a highly sparse change in the representation. Disentangled representations offer several advantages – (i) Invariance: it is easier to derive representations that are invariant to nuisance factors by simply marginalizing over the corresponding dimensions, (ii) Transferability: they are arguably more suitable for transfer learning as most of the key underlying generative factors appear segregated along feature dimensions, (iii) Interpretability: a human expert may be able to assign meanings to the dimensions, (iv) Conditioning and intervention: they allow for interpretable conditioning and/or intervention over a subset of the latents and observe the effects on other nodes in the graph. Indeed, the importance of learning disentangled representations has been argued in several recent works (Bengio et al., 2013; Lake et al., 2016; Ridgeway, 2016).

Recognizing the significance of disentangled representations, several attempts have been made in this direction in the past (Ridgeway, 2016). Much of the earlier work assumes some sort of supervision in terms of: (i) partial or full access to the generative factors per instance (Reed et al., 2014; Yang et al., 2015; Kulkarni et al., 2015; Karaletsos et al., 2015), (ii) knowledge about the nature of generative factors (e.g, translation, rotation, etc.) (Hinton et al., 2011; Cohen & Welling, 2014), (iii) knowledge about the changes in the generative factors across observations (e.g., sparse changes in consecutive frames of a Video) (Goroshin et al., 2015; Whitney et al., 2016; Fraccaro et al., 2017; Denton & Birodkar, 2017; Hsu et al., 2017), (iv) knowledge of a complementary signal to infer representations that are conditionally independent of itThe representation itself can still be entangled in rest of the generative factors. (Cheung et al., 2014; Mathieu et al., 2016; Siddharth et al., 2017). However, in most real scenarios, we only have access to raw observations without any supervision about the generative factors. It is a challenging problem and many of the earlier attempts have not been able to scale well for realistic settings (Schmidhuber, 1992; Desjardins et al., 2012; Cohen & Welling, 2015) (see also, Higgins et al. (2017)).

Recently, Chen et al. (2016) proposed an approach to learn a generative model with disentangled factors based on Generative Adversarial Networks (GAN) (Goodfellow et al., 2014), however implicit generative models like GANs lack an effective inference mechanismThere have been a few recent attempts in this direction for visual data (Dumoulin et al., 2016; Donahue et al., 2016; Kumar et al., 2017) but often the reconstructed samples are semantically quite far from the input samples, sometimes even changing in the object classes., which hinders its applicability to the problem of learning disentangled representations. More recently, Higgins et al. (2017) proposed an approach based on Variational AutoEncoder (VAE) Kingma & Welling (2013) for inferring disentangled factors. The inferred latents using their method (termed as β\beta-VAE ) are empirically shown to have better disentangling properties, however the method deviates from the basic principles of variational inference, creating increased tension between observed data likelihood and disentanglement. This in turn leads to poor quality of generated samples as observed in (Higgins et al., 2017).

In this work, we propose a principled approach for inference of disentangled latent factors based on the popular and scalable framework of amortized variational inference (Kingma & Welling, 2013; Stuhlmüller et al., 2013; Gershman & Goodman, 2014; Rezende et al., 2014) powered by stochastic optimization (Hoffman et al., 2013; Kingma & Welling, 2013; Rezende et al., 2014). Disentanglement is encouraged by introducing a regularizer over the induced inferred prior. Unlike β\beta-VAE (Higgins et al., 2017), our approach does not introduce any extra conflict between disentanglement of the latents and the observed data likelihood, which is reflected in the overall quality of the generated samples that matches the VAE and is much better than β\beta-VAE. This does not come at the cost of higher entanglement and our approach also outperforms β\beta-VAE in disentangling the latents as measured by various quantitative metrics. We also propose a new disentanglement metric, called Separated Attribute Predictability or SAP, which is better aligned with the qualitative disentanglement observed in the decoder’s output compared to the existing metrics.

Formulation

The ELBO (the objective at the right side of Eq. 1) lower bounds the log-likelihood of observed data, and the gap vanishes at the global optimum. Often, the density forms of p(z)p(\mathbf{z}) and qϕ(z∣x)q_{\phi}(\mathbf{z}|\mathbf{x}) are chosen such that their KL-divergence can be written analytically in a closed-form expression (e.g., p(z)p(\mathbf{z}) is N(0,I)N(0,I) and qϕ(z∣x)q_{\phi}(\mathbf{z}|\mathbf{x}) is N(μϕ(x),Σϕ(x))N(\mu_{\phi}(\mathbf{x}),\Sigma_{\phi}(\mathbf{x}))) (Kingma & Welling, 2013). In such cases, the ELBO can be efficiently optimized (to a stationary point) using stochastic first order methods where both expectations are estimated using mini-batches. Further, in cases when qϕ(⋅)q_{\phi}(\cdot) can be written as a continuous transformation of a fixed base distribution (e.g., the standard normal distribution), a low variance estimate of the gradient over ϕ\phi can be obtained by coordinate transformation (also referred as reparametrization) (Fu, 2006; Kingma & Welling, 2013; Rezende et al., 2014).

Most VAE based generative models for real datasets (e.g., text, images, etc.) already work with a relatively simple and disentangled prior p(z)p(\mathbf{z}) having no interaction among the latent dimensions (e.g., the standard Gaussian N(0,I)N(0,I)) (Bowman et al., 2015; Miao et al., 2016; Hou et al., 2017; Zhao et al., 2017). The complexity of the observed data is absorbed in the conditional distribution pθ(x∣z)p_{\theta}(\mathbf{x}|\mathbf{z}) which encodes the interactions among the latents. Hence, as far as the generative modeling is concerned, disentangled prior sets us in the right direction.

2 Inferring disentangled latents

Although the generative model starts with a disentangled prior, our main objective is to infer disentangled latents which are potentially conducive for various goals mentioned in Sec. 1 (e.g., invariance, transferability, interpretability). To this end, we consider the density over the inferred latents induced by the approximate posterior inference mechanism,

which we will subsequently refer to as the inferred prior or expected variational posterior (p(x)p(\mathbf{x}) is the true data distribution that we have only samples from). For inferring disentangled factors, this should be factorizable along the dimensions, i.e., qϕ(z)=∏iqi(zi)q_{\phi}(\mathbf{z})=\prod_{i}q_{i}(z_{i}), or equivalently qi∣j(zi∣zj)=qi(zi), ∀ i,jq_{i|j}(z_{i}|z_{j})=q_{i}(z_{i}),\,\forall\,i,j. This can be achieved by minimizing a suitable distance between the inferred prior qϕ(z)q_{\phi}(\mathbf{z}) and the disentangled generative prior p(z)p(\mathbf{z}). We can also define expected posterior as pθ(z)=∫pθ(z∣x)p(x)dxp_{\theta}(\mathbf{z})=\int p_{\theta}(\mathbf{z}|\mathbf{x})p(\mathbf{x})d\mathbf{x}. If we take KL-divergence as our choice of distance, by relying on its pairwise convexity (i.e., KL(λp1+(1−λ)p2∥λq1+(1−λ)q2)≤λKL(p1∥q1)+(1−λ)KL(p2∥q2)\textrm{KL}(\lambda p_{1}+(1-\lambda)p_{2}\|\lambda q_{1}+(1-\lambda)q_{2})\leq\lambda\textrm{KL}(p_{1}\|q_{1})+(1-\lambda)\textrm{KL}(p_{2}\|q_{2})) (Van Erven & Harremos, 2014), we can show that the distance between qϕ(z)q_{\phi}(\mathbf{z}) and pθ(z)p_{\theta}(\mathbf{z}) is bounded by the objective of the variational inference:

where λ\lambda controls its contribution to the overall objective. We refer to this as DIP-VAE (for Disentangled Inferred Prior) subsequently.

Optimizing (4) directly is not tractable if D(⋅,⋅)D(\cdot,\cdot) is taken to be the KL-divergence KL(qϕ(z)∥p(z))\textrm{KL}(q_{\phi}(\mathbf{z})\|p(\mathbf{z})), which does not have a closed-form expression. One possibility is use the variational formulation of the KL-divergence (Nguyen et al., 2010; Nowozin et al., 2016) that needs only samples from qϕ(z)q_{\phi}(\mathbf{z}) and p(z)p(\mathbf{z}) to estimate a lower bound to KL(qϕ(z)∥p(z))\textrm{KL}(q_{\phi}(\mathbf{z})\|p(\mathbf{z})). However, this would involve optimizing for a third set of parameters ψ\psi for the KL-divergence estimator, and would also change the optimization to a saddle-point (min-max) problem which has its own optimization challenges (e.g., gradient vanishing as encountered in training generative adversarial networks with KL or Jensen-Shannon (JS) divergences (Goodfellow et al., 2014; Arjovsky & Bottou, 2017)). Taking DD to be another suitable distance between qϕ(z)q_{\phi}(\mathbf{z}) and p(z)p(\mathbf{z}) (e.g., integral probability metrics like Wasserstein distance (Sriperumbudur et al., 2009)) might alleviate some of these issues (Arjovsky et al., 2017) but will still involve complicating the optimization to a saddle point problem in three set of parametersNonparametric distances like maximum mean discrepancy (MMD) with a characteristic kernel (Gretton et al., 2012) is also an option, however it has its own challenges when combined with stochastic optimization (Dziugaite et al., 2015; Li et al., 2015).. It should also be noted that using these variational forms of the distances will still leave us with an approximation to the actual distance.

The regularization terms involving Covp(x)[μϕ(x)]\text{Cov}_{p(\mathbf{x})}[\bm{\mu}_{\phi}(\mathbf{x})] in the above objective (6) can be efficiently optimized using SGD, where Covp(x)[μϕ(x)]\text{Cov}_{p(\mathbf{x})}[\bm{\mu}_{\phi}(\mathbf{x})] can be estimated using the current minibatchWe also tried an alternative of maintaining a running estimate of Covp(x)[μϕ(x)]\text{Cov}_{p(\mathbf{x})}[\bm{\mu}_{\phi}(\mathbf{x})] which is updated with every minibatch of x∼p(x)\mathbf{x}\sim p(\mathbf{x}), however we did not observe a significant improvement over the simpler approach of estimating these using only current minibatch..

For DIP-VAE-II, we have the following optimization problem:

3 Comparison with β𝛽\beta-VAE

Recently proposed β\beta-VAE (Higgins et al., 2017) proposes to modify the ELBO by upweighting the KL(qϕ(z∣x)∥p(z))\textrm{KL}(q_{\phi}(\mathbf{z}|\mathbf{x})\|p(\mathbf{z})) term in order to encourage the inference of disentangled factors:

where β\beta is taken to be great than 11. Higher β\beta is argued to encourage disentanglement at the cost of reconstruction error (the likelihood term in the ELBO). Authors report empirical results with β\beta ranging from 44 to 250250 depending on the dataset. As already mentioned, most VAE models proposed in the literature, including β\beta-VAE, work with N(0,I)N(\mathbf{0},\mathbf{I}) as the prior p(z)p(\mathbf{z}) and N(μϕ(x),Σϕ(x))N(\bm{\mu}_{\phi}(\mathbf{x}),\bm{\Sigma}_{\phi}(\mathbf{x})) with diagonal Σϕ(x)\bm{\Sigma}_{\phi}(\mathbf{x}) as the approximate posterior qϕ(z∣x)q_{\phi}(\mathbf{z}|\mathbf{x}). This reduces the objective (8) to

For high values of β\beta, β\beta-VAE would try to pull μϕ(x)\bm{\mu}_{\phi}(\mathbf{x}) towards zero and Σϕ(x)\bm{\Sigma}_{\phi}(\mathbf{x}) towards the identity matrix (as the minimum of x−ln⁡xx-\ln x for x>0x>0 is at x=1x=1), thus making the approximate posterior qϕ(z∣x)q_{\phi}(\mathbf{z}|\mathbf{x}) insensitive to the observations. This is also reflected in the quality of the reconstructed samples which is worse than VAE (β=1\beta=1), particularly for high values of β\beta. Our proposed method does not have such increased tension between the likelihood term and the disentanglement objective, and the sample quality with our method is on par with the VAE.

Quantifying disentanglement: SAP Score

Higgins et al. (2017) propose a metric to evaluate the disentanglement performance of the inference mechanism, assuming that the ground truth generative factors are available. It works by first sampling a generative factor yy, followed by sampling LL pairs of examples such that for each pair, the sampled generative factor takes the same value. Given the inferred zx:=μϕ(x)\mathbf{z}_{x}:=\bm{\mu}_{\phi}(\mathbf{x}) for each example x\mathbf{x}, they compute the absolute difference of these vectors for each pair, followed by averaging these difference vectors. This average difference vector is assigned the label of yy. By sampling nn such minibatches of LL pairs, we get nn such averaged difference vectors for the factor yy. This process is repeated for all generative factors. A low capacity multiclass classifier is then trained on these vectors to predict the identities of the corresponding generative factors. Accuracy of this classifier on the difference vectors for test set is taken to be a measure of disentanglement. We evaluate the proposed method on this metric and refer to this as Z-diff score subsequently.

We observe in our experiments that the Z-diff score (Higgins et al., 2017) is not correlated well with the qualitative disentanglement at the decoder’s output as seen in the latent traversal plots (obtained by varying only one latent while keeping the other latents fixed). It also depends on the multiclass classifier used to obtain the score. We propose a new metric, referred as Separated Attribute Predictability (SAP) score, that is better aligned with the qualitative disentanglement observed in the latent traversals and also does not involve training any classifier. It is computed as follows: (i) We first construct a d×kd\times k score matrix SS (for dd latents and kk generative factors) whose ijij’th entry is the linear regression or classification score (depending on the generative factor type) of predicting jj’th factor using only ii’th latent [μϕ(x)]i[\bm{\mu}_{\phi}(\mathbf{x})]_{i}. For regression, we take this to be the R2R^{2} score obtained with fitting a line (slope and intercept) that minimizes the linear regression error (for the test examples). The R2R^{2} score is given by (Cov([μϕ(x)]i,yj)σ[μϕ(x)]iσyj)2\left(\frac{\text{Cov}([\bm{\mu}_{\phi}(\mathbf{x})]_{i},\mathbf{y}_{j})}{\sigma_{[\bm{\mu}_{\phi}(\mathbf{x})]_{i}}\sigma_{\mathbf{y}_{j}}}\right)^{2} and ranges from to 11, with a score of 11 indicating that a linear function of the ii’th inferred latent explains all variability in the jj’th generative factor. For classification, we fit one or more thresholds (real numbers) directly on ii’th inferred latents for the test examples that minimize the balanced classification errors, and take SijS_{i}j to be the balanced classification accuracy of the jj’th generative factor. For inactive latent dimensions (having σ[μϕ(x)]i=[Covp(x)[μϕ(x)]]ii\sigma_{[\bm{\mu}_{\phi}(\mathbf{x})]_{i}}=[\text{Cov}_{p(x)}[\bm{\mu}_{\phi}(\mathbf{x})]]_{ii} close to ), we take SijS_{ij} to be . (ii) For each column of the score matrix SS which corresponds to a generative factor, we take the difference of top two entries (corresponding to top two most predictive latent dimensions), and then take the mean of these differences as the final SAP score. Considering just the top scoring latent dimension for each generative factor is not enough as it does not rule out the possibility of the factor being captured by other latents. A high SAP score indicates that each generative factor is primarily captured in only one latent dimension. Note that a high SAP score does not rule out one latent dimension capturing two or more generative factors well, however in many cases this would be due to the generative factors themselves being correlated with each other, which can be verified empirically using ground truth values of the generative factors (when available). Further, a low SAP score does not rule out good disentanglement in cases when two (or more) latent dimensions might be correlated strongly with the same generative factor and poorly with other generative factors. The generated examples using single latent traversals may not be realistic for such models, and DIP-VAE discourages this from happening by enforcing decorrelation of the latents. However, the SAP score computation can be adapted to such cases by grouping the latent dimensions based on correlations and getting the score matrix at group level, which can be fed as input to the second step to get the final SAP score.

Experiments

We evaluate our proposed method, DIP-VAE, on three datasets – (i) CelebA (Liu et al., 2015): It consists of 202,599202,599 RGB face images of celebrities. We use 64×64×364\times 64\times 3 cropped images as used in several earlier works, using 90%90\% for training and 10%10\% for test. (ii) 3D Chairs (Aubry et al., 2014): It consists of 1393 chair CAD models, with each model rendered from 31 azimuth angles and 2 elevation angles. Following earlier work (Yang et al., 2015; Dosovitskiy et al., 2015) that ignores near-duplicates, we use a subset of 809 chair models in our experiments. We use the binary masks of the chairs as the observed data in our experiments following (Higgins et al., 2017). First 80%80\% of the models are used for training and the rest are used for test. (iii) 2D Shapes (Matthey et al., 2017): This is a synthetic dataset of binary 2D shapes generated from the Cartesian product of the shape (heart, oval and square), xx-position (32 values), yy-position (32 values), scale (6 values) and rotation (40 values). We consider two baselines for the task of unsupervised inference of disentangled factors: (i) VAE (Kingma & Welling, 2013; Rezende et al., 2014), and (ii) the recently proposed β\beta-VAE (Higgins et al., 2017). To be consistent with the evaluations in (Higgins et al., 2017), we use the same CNN network architectures (for our encoder and decoder), and same latent dimensions as used in (Higgins et al., 2017) for CelebA, 3D Chairs, 2D Shapes datasets.

Disentanglement scores and reconstruction error. For the Z-diff score (Higgins et al., 2017), in all our experiments we use a one-vs-rest linear SVM with weight on the hinge loss CC set to 0.010.01 and weight on the regularizer set to 11. Table 1 shows the Z-diff scores and the proposed SAP scores along with reconstruction error (which directly corresponds to the data likelihood) for the test sets of CelebA and 2D Shapes data. Further we also show the plots of how the Z-diff score and the proposed SAP score change with the reconstruction error as we vary the hyperparameter for both methods (β\beta and λod\lambda_{od}, respectively) in Fig. 1 (for 2D Shapes data) and Fig. 2 (for CelebA data). The proposed DIP-VAE-I gives much higher Z-diff score at little to no cost on the reconstruction error when compared with VAE (β=1\beta=1) and β\beta-VAE, for both 2D Shapes and CelebA datasets. However, we observe in the decoder’s output for single latent traversals (varying a single latent while keeping others fixed, shown in Fig. 3 and Fig. 4) that a high Z-diff score is not necessarily a good indicator of disentanglement. Indeed, for 2D Shapes data, DIP-VAE-I has a higher Z-diff score (98.798.7) and almost an order of magnitude lower reconstruction error than β\beta-VAE for β=60\beta=60, however comparing the latent traversals of β\beta-VAE in Fig. 3 and DIP-VAE-I in Fig. 4 indicate a better disentanglement for β\beta-VAE for β=60\beta=60 (though at the cost of much worse reconstruction where every generated sample looks like a hazy blob). On the other hand, we find the proposed SAP score to be correlated well with the qualitative disentanglement seen in the latent traversal plots. This is reflected in the higher SAP score of β\beta-VAE for β=60\beta=60 than DIP-VAE-I. We also observe that for 2D Shapes data, DIP-VAE-II gives a much better trade-off between disentanglement (measured by the SAP score) and reconstruction error than both DIP-VAE-I and β\beta-VAE, as shown quantitatively in Fig. 1 and qualitatively in the latent traversal plots in Fig. 3. The reason is that DIP-VAE-I enforces [Covp(x)[μϕ(x)]]ii[\text{Cov}_{p(x)}[\bm{\mu}_{\phi}(\mathbf{x})]]_{ii} to be close to 11 and this may affect the disentanglement adversely by splitting a generative factor across multiple latents for 2D Shapes where the generative factors are much less than the latent dimension. For real datasets having lots of factors with complex generative processes, such as CelebA, DIP-VAE-I is expected to work well which can be seen in Fig. 2 where DIP-AVE-I yields a much lower reconstruction error with a higher SAP score (as well as higher Z-diff scores).

Binary attribute classification for CelebA. We also experiment with predicting the binary attribute values for each test example in CelebA from the inferred μϕ(x)\bm{\mu}_{\phi}(\mathbf{x}). For each attribute kk, we compute the attribute vector wk=1∣xi:yik=1∣∑xi:yik=1μϕ(xi)−1∣xi:yik=0∣∑xi:yik=0μϕ(xi)\mathbf{w}^{k}=\frac{1}{|\mathbf{x}_{i}:y_{i}^{k}=1|}\sum_{\mathbf{x}_{i}:y_{i}^{k}=1}\mu_{\phi}(\mathbf{x}_{i})-\frac{1}{|\mathbf{x}_{i}:y_{i}^{k}=0|}\sum_{\mathbf{x}_{i}:y_{i}^{k}=0}\mu_{\phi}(\mathbf{x}_{i}) from the training set, and project the μϕ(x)\bm{\mu}_{\phi}(\mathbf{x}) along these vectors. A bias is learned on these scalars (by minimizing hinge loss) which is then used for classifying the test examples. Table 4 shows the results for the attribute which show the highest change across various methods (most other attribute accuracies do not change). The proposed DIP-VAE outperforms both VAE and β\beta-VAE for most attributes. The performance of β\beta-VAE gets worse as β\beta is increased further.

Invariance and Equivariance. Disentanglement is closely connected to invariance and equivariance of representations. If R:x→zR:\mathbf{x}\to\mathbf{z} is a function that maps the observations to the feature representions, equivariance (with respect to TT) implies that a primitive transformation TT of the input results in a corresponding transformation T′T^{\prime} of the feature, i.e., R(T(x))=T′(R(x))R(T(\mathbf{x}))=T^{\prime}(R(\mathbf{x})). Disentanglement requires that T′T^{\prime} acts only on a small subset of dimensions of R(x)R(\mathbf{x}) (a sparse action). In this sense, equivariance is a more general notion encompassing disentanglement as a special case, however this special case carries additional benefits of interpretability, ease of transferrability, etc. Invariance is also a special case of equivariance which requires T′T^{\prime} to be identity for RR to be invariant to the action of TT on the input observations. However, invariance can obtained more easily from disentangled representations than from equivariant representations by simply marginalizing the appropriate subset of dimensions. There exists a lot of prior work in the literature on equivariant and invariant feature learning, mostly under the supervised setting which assumes the knowledge about the nature of input transformations (e.g., rotations, translations, scaling for images, etc.) (Schmidt & Roth, 2012; Bruna & Mallat, 2013; Anselmi et al., 2014; 2016; Cohen & Welling, 2016; Dieleman et al., 2016; Haasdonk et al., 2005; Mroueh et al., 2015; Raj et al., 2017).

We proposed a principled variational framework to infer disentangled latents from unlabeled observations. Unlike β\beta-VAE, our variational objective does not have any conflict between the data log-likelihood and the disentanglement of the inferred latents, which is reflected in the empirical results. We also proposed the SAP disentanglement metric that is much better correlated with the qualitative disentanglement seen in the latent traversals than the Z-diff score Higgins et al. (2017). An interesting direction for future work is to take into account the sampling biases in the generative process, both natural (e.g., sampling the female gender makes it unlikely to sample beard for face images in CelebA) as well as artificial (e.g., a collection of face images that contain much more smiling faces for males than females misleading us to believe p(gender,smile)≠p(gender)p(smile)p(\text{gender,smile})\neq p(\text{gender})p(\text{smile})), which makes the problem challenging and also somewhat less well defined (at least in the case of natural biases). Effective use of disentangled representations for transfer learning is another interesting direction for future work.

Appendix A Latent traversals for 2D Shapes and Chairs dataset