Consistency Regularization for Variational Auto-Encoders
Samarth Sinha, Adji B. Dieng
Introduction
Variational auto-encoders (vaes) have significantly impacted research on unsupervised learning. They have been used in several areas, including density estimation (Kingma & Welling, 2013; Rezende et al., 2014), image generation (Gregor et al., 2015), text generation (Bowman et al., 2015; Fang et al., 2019), music generation (Roberts et al., 2018), topic modeling (Miao et al., 2016; Dieng et al., 2019), and recommendation systems (Liang et al., 2018). Vaes have also been used for different representation learning problems such as semi-supervised learning (Kingma et al., 2014), anomaly detection (An & Cho, 2015; Zimmerer et al., 2018), language modeling Bowman et al. (2015), active learning (Sinha et al., 2019), continual learning (Achille et al., 2018), and motion prediction of agents (Walker et al., 2016). This widespread application of vae representations makes it critical that we focus on improving them.
vaes extend deterministic auto-encoders to probabilistic generative modeling. The encoder of a vae parameterizes an approximate posterior distribution over latent variables of a generative model. The encoder is shared between all observations, which amortizes the cost of posterior inference. Once fitted, the encoder of a vae can be used to obtain low-dimensional representations of data, (e.g. for downstream tasks.) The quality of these representations is therefore very important to a successful application of vaes.
Researchers have looked at ways to improve the quality of the latent representations of vaes, often tackling the so-called latent variable collapse problem—in which the approximate posterior distribution induced by the encoder collapses to the prior over the latent variables (Bowman et al., 2015; Kim et al., 2018; Dieng et al., 2018; He et al., 2019; Fu et al., 2019).
In this paper, we focus on a different problem pertaining to the latent representations of vaes for image data. Indeed, the encoder of a fitted vae tends to map an image and a semantics-preserving transformation of that image to different parts in the latent space. This “inconsistency" of the encoder affects the quality of the learned representations and generalization. We propose a method to enforce consistency in vaes. The idea is simple and consists in maximizing the likelihood of the images while minimizing the Kullback-Leibler (kl) divergence between the approximate posterior distribution induced by the encoder when conditioning on the image, on one hand, and its transformation, on the other hand. This regularization technique can be applied to any vae variant to improve the quality of the learned representations and boost generalization performance. We call a vae with this form of regularization, a consistency-regularized variational auto-encoder (cr-vae).
Figure 1 illustrates the inconsistency problem of vaes and how cr-vaes address this problem on mnist. The red dots are representations of a few images and the blue dots are the representations of their transformations. We applied semantics-preserving transformations: rotation, translation, and scaling. The vae maps each image and its transformation to different parts in the latent space as evidenced by the long arrows connecting each pair (a). Even when we include the transformed images to the data and fit the vae the inconsistency problem still occurs (b). The cr-vae does not suffer from the inconsistency problem; it maps each image and its transformation to nearby areas in the latent space, as evidenced by the short arrows connecting each pair (c).
In our experiments (see Section 4), we apply the proposed technique to four vae variants, the original vae (Kingma & Welling, 2013), the importance-weighted auto-encoder (iwae) (Burda et al., 2015), the -vae (Higgins et al., 2017), and the nouveau variational auto-encoder (nvae) (Vahdat & Kautz, 2020). We found, on four different benchmark datasets, that cr-vaes always yield better representations and generalize better than their base vaes. In particular, consistency-regularized nouveau variational auto-encoders (cr-nvaes) yield state-of-the-art performance on mnist and cifar-10. We also applied cr-vaes to 3D data where these conclusions still hold.
Method
We consider a latent-variable model , where denotes an observation and is its associated latent variable. The marginal is a prior over the latent variable and is an exponential family distribution whose natural parameter is a function of parameterized by , e.g. through a neural network. Our goal is to learn the parameters and a posterior distribution over the latent variables. The approach of vaes is to maximize the evidence lower bound (elbo), a lower bound on the log marginal likelihood of the data,
where is an approximate posterior distribution over the latent variables. The idea of a vae is to let the parameters of the distribution be given by the output of a neural network, with parameters , that takes as input. The parameters and are then jointly optimized by maximizing a Monte Carlo approximation of the elbo using the reparameterization trick (Kingma & Welling, 2013).
We now propose a regularization method that ensures consistency of the encoder of a vae. We call a vae with such a regularization a cr-vae. The regularization proposed is applicable to many variants of the vae such as the iwae (Burda et al., 2015), the -vae (Higgins et al., 2017), and the nvae (Vahdat & Kautz, 2020). In what follows, we use the standard vae, the one that maximizes Equation 1, as the base vae to regularize to illustrate the method.
Here is a semantics-preserving transformation of the image , e.g. translation with random length drawn from for some threshold . A cr-vae then maximizes
where the regularization term is
Maximizing the objective in Equation 3 maximizes the likelihood of the data and their augmentations while enforcing consistency through . Minimizing , which only affects the encoder (with parameters ), forces each observation and the corresponding augmentations to lie close to each other in the latent space. The hyperparameter controls the strength of this constraint.
Related Work
More recently, consistency regularization has been applied to generative adversarial networks (gans) (Goodfellow et al., 2014). Indeed Wei et al. (2018) and Zhang et al. (2020) show that applying consistency regularization on the discriminator of a gan—also a classifier—can substantially improve its performance.
The idea we develop in this paper differs from the works above in two ways. First, it applies consistency regularization to vaes for image data. Second, it leverages consistency regularization, not in the label or logit space, as done in the works mentioned above, but in the latent space.
Although different, consistency regularization for vaes relates to works that study ways to constrain the sensitivity of encoders to various perturbations. For example, denoising auto-encoders (daes) and their variants (Vincent et al., 2008, 2010) corrupt an image into , typically using Gaussian noise, and then minimize the distance between the reconstruction of and the un-corrupted image . The motivation is to learn representations that are insensitive to the added noise. Our work differs in that we do not constrain the decoder to recover the original image from the corrupted image but, rather, to constrain the encoder to recover the latent representation of the original image from the corrupted image via a kl divergence minimization constraint.
Contractive auto-encoders (caes) (Rifai et al., 2011) share a similar goal with cr-vaes. A cae is an auto-encoder whose encoder is constrained by minimizing the norm of the Jacobian of the output of the encoder with respect to the input image. This norm constraint on the Jacobian forces the representations learned by the encoder to be insensitive to changes in the input. Our work differs in several main ways. First, cr-vaes are not deterministic auto-encoders, contrary to caes. We can easily sample from a cr-vae, as for any vae, which is not the case for a cae. Second, a cae does not apply transformations to the input image, which limits the sensitivities it can learn to limit to those exhibited in the training set. Finally, caes use the Jacobian to impose a consistency constraint, which are not as easy to compute as the kl divergence we use on the variational distribution induced by the encoder.
Empirical Study
In this section we show that a cr-vae improves the learned representations of its base vae and positively affects generalization performance We also show that the proposed regularization method is amenable to different vae variants by applying it not only to the original vae but also to the iwae, the -vae, and the nvae. We showcase the importance of the KL regularization term by conducting an ablation study. We found that only regularizing with data augmentation improves performance but that accounting for the kl term () further improves the quality of the learned representations and generalization.
We will conduct three sets of experiments. In the first experiment, we will apply the regularization method proposed in this paper to standard vaes such as the original vae, the iwae, and the -vae. We use mnist, omniglot, and celeba as datasets for this experiment. For celeba, we choose the x resolution for this experiment. Our results show that adding consistency regularization always improves upon the base vae, both in terms of the quality of the learned representations and generalization. We conduct an ablation study and also report performance of the different vae variants above when they are fitted with the original data and their augmentations. The results from this ablation highlight the importance of setting .
In the second set of experiments we apply our method to a large-scale vae, the latest nvae (Vahdat & Kautz, 2020). We use mnist, cifar-10, and celeba as datasets for this experiment. We increased the resolution for the celeba dataset for this experiment to x. We reach the same conclusions as for the first sets of experiments; cr-vaes improve the learned representations and generalization of their base vaes. In this particular setting, the cr-nvae achieves state-of-the-art generalization performance on both mnist and cifar-10. This state-of-the-art performance couldn’t be reach simply by training the nvae with augmentations, as our results show.
Finally, in a third set of experiments, we apply our regularization technique to a 3D point-cloud dataset called ShapeNet (Chang et al., 2015). We adapt a high-performing auto-encoding method called FoldingNet (Yang et al., 2018) to its vae counterpart and apply the method we described in this paper to that vae variant on the ShapeNet dataset. We found that adding consistency regularization yields better learned representations.
We next describe in great detail the set up for each of these experiments and the results showcasing the usefulness of the regularization method we propose in this paper.
We apply consistency regularization, as described in this paper, to the original vae, the iwae, and the -vae. We now describe the set up and results for this experiment.
Datasets. We study three benchmark datasets that we briefly describe below. We first consider mnist. mnist is a handwritten digit recognition dataset with images in the training set and images in the test set (LeCun, 1998). We form a validation set of images randomly sampled from the training set.
We also consider omniglot, a handwritten alphabet recognition dataset (Lake et al., 2011). This dataset is composed of images. We use randomly sampled images for training and for validation and the remaining samples for testing.
Finally we consider celeba. It is a dataset of faces, consisting of images for training, images for validation, and images for testing (Liu et al., 2018). We set the resolution to x for this experiment.
Evaluation metrics. The regularization method we propose in this paper is mainly aimed at improving the learned representations of vaes. To assess these representations we use three metrics: mutual information, number of active latent units, and accuracy on a downstream classification task. We also evaluate the effect of the proposed method on generalization to unseen data. For that we also report negative log-likelihood. We define each of these metrics next.
Mutual information (MI). The first quality metric is the mutual information between the observations and the latents under the joint distribution induced by the encoder,
where is the empirical data distribution and is the aggregated posterior, the marginal over induced by the joint distribution defined by and . The mutual information is intractable but we can approximate it with Monte Carlo. Higher mutual information corresponds to more interpretable latent variables.
Number of active latent units (AU). The second quality metrics we consider is the number of active latent units (AU). It is defined in Burda et al. (2015) and measures the “activity" of a dimension of the latent variables . A latent dimension is “active" if
where is a threshold defined by the user. For our experiments we set . The higher the number of latent active units, the better the learned representations.
Accuracy on downstream classification. This metric is calculated by fitting a given vae, taking the learned representations for each data in the test set and computing the accuracy from the prediction of the labels of the images in that same test set by a classifier fitted on the training set. This metric is only applicable to labelled datasets.
Negative log-likelihood. We use negative held-out log-likelihood to assess generalization. Consider an unseen data , its negative held-out log-likelihood under the fitted model is
This is intractable and we approximate it using Monte Carlo,
where .
Settings. The vaes are built on the same architecture as Tolstikhin et al. (2017). The networks are trained with the Adam optimizer with a learning rate of (Kingma & Ba, 2014) and trained for epochs with a batch size of . We set the dimensionality of the latent variables to , therefore the maximum number of active latent units in the latent space is . We found to be best according to cross-validation using held-out log-likelihood and exploring the range datasets. In an ablation study we explore . For the -vae we set and study both and , two regimes under which the -vae performs qualitatively very differently (Higgins et al., 2017). All experiments were done on a GPU cluster consisting of Nvidia P100 and RTX. The training took approximately 1 day for most experiments.
Results. Table 1 shows that on all the three benchmark datasets all the different vae variants we studied, consistency regularization as developed in this paper always improves the quality of the learned representations as measured by mutual information and the number of active latent units. These results are confirmed by the numbers shown in Table 2 where cr-vaes always lead to better accuracy on downstream classification.
We proposed consistency regularization as a way to improve the quality of the learned representations. Incidentally, Table 3 also shows that it can improve generalization as measured by negative log-likelihood.
Ablation Study. We now look at the impact of each factor that goes into the regularization method we introduced in this paper using mnist. We test the impact of the regularization term and the impact of the choice of augmentation on all metrics. Table 4 and Table 5 show the results.
Table 4 shows that even small consistency regularization (a small value) results in improvement over the base vae but that a large enough value can hurt performance.
Table 5 shows that rotations and translations are more important than scaling, but the combination of all three augmentations works best for cr-vaes.
Comparison to Contrastive Learning. We look at how cr-vaes compare against a popular and advanced contrastive-learning-based technique, the triplet loss (Schroff et al., 2015) using mnist. Table 6 shows that the cr-vae outperforms the triplet loss on both generalization performance and quality of learned representations. Table 6 also confirms existing literature showing simply applying augmentations can outperform complex contrastive learning-based methods such as the triplet loss (Kostrikov et al., 2020; Sinha & Garg, 2021).
2 Application to the large-scale nvae on benchmark datasets
Along with standard VAE variants, we also experiment with a large scale state-of-the-art vae, the nvae (Vahdat & Kautz, 2020). Similar to before, we simply add consistency regularization using the image-based augmentations techniques to the NVAE model and experiment on benchmark datasets: mnist (LeCun, 1998), cifar-10 (Krizhevsky et al., 2009) and celeba (Liu et al., 2018).
The results for large scale generative modeling are tabulated in Table 8 and Table 7, where we see that using cr-nvae we are able to learn representations that yield better accuracy on downstream classification and set new state-of-the-art values on each of the datasets, improving upon the baseline log-likelihood values. This shows the ability of consistency regularization to work at scale on challenging generative modeling tasks.
3 Application to the FoldingNet on 3D point-cloud data
Along with working with image data, we additionally experiment with 3D point cloud data using a FoldingNet Yang et al. (2018) and the ShapeNet dataset Chang et al. (2015) which consists of 55 distinct object classes. FoldingNet learns a deep AutoEncoder to learn unsupervised representations from the point cloud data. To add consistency regularization, we first substitute the AutoEncoder to a vae by adding the KL term from the ELBO to the baseline FoldingNet. We then add the additional consistency regularization KL term to the latent space of FoldingNet.
For the ShapeNet point cloud data, we perform data augmentation using a similar scheme to what we did for the previous experiments, we randomly translate, rotate and add jitter to the coordinates of the point cloud data. We follow the same scheme detailed in FoldingNet (Yang et al., 2018).
We train both the FoldingNet turned in a vae and the CR-FoldingNet with these augmentations. To train CR-FoldingNet, we additionally apply the consistency regularization term as proposed in Equation 3. The results on the validation set for reconstruction (as measured by Chamfer distance) and accuracy are shown in Table 9.
We also visualize the point clouds reconstructions and interpolations between 3 different object classes using a CR-FoldingNet in Figure 2. We perform 4 interpolation steps for each of the objects, to highlight the interpretable learned latent space. Additionally, we perform the same interpolation on the baseline FoldingNet model. We show these interpolations in the appendix.
Conclusion
We proposed a simple regularization technique to constrain encoders of vaes to learn similar latent representations for an image and a semantics-preserving transformation of the image. The idea consists in maximizing the likelihood of the pair of images while minimizing the kl divergence between the variational distribution induced by the encoder when conditioning on the image on one hand, and its transformation, on the other hand. We applied this technique to several vae variants on several datasets, including a 3D dataset. We found it always leads to better learned representations and also better generalization to unseen data. In particular, when applied to the nvae, the regularization technique we developed in this paper yields state-of-the-art results on mnist and cifar-10.
Broader Impact
In this paper, we propose a simple method that performs a KL-based consistency regularization scheme using data augmentation for vaes. The broader impact of the study includes practical applications such as graphics and computer vision applications. The method we propose improves the learned representations of vaes, and as an artifact, also improves their generalization to unseen data. In this regard, any implications of vaes also apply to this work. For example, the generative model fit by a vae may be used to generate artificial data such as images, text, and 3D objects. Biases may arise as a result of poor data selection. Furthermore, text generated from generative systems may amplify harmful speech contained in the data. However, the method we propose can also improve the performance of vaes when used in certain practical domains as we discussed in the introduction of the paper.
Acknowledgements
We thank Kevin Murphy, Ben Poole, and Augustus Odena for their comments on this work.