Learning Latent Space Energy-Based Prior Model
Bo Pang, Tian Han, Erik Nijkamp, Song-Chun Zhu, Ying Nian Wu
Introduction
In recent years, deep generative models have achieved impressive successes in image and text generation. A particularly simple and powerful model is the generator model , which assumes that the observed example is generated by a low-dimensional latent vector via a top-down network, and the latent vector follows a non-informative prior distribution, such as uniform or isotropic Gaussian distribution. While we can learn an expressive top-down network to map the prior distribution to the data distribution, we can also learn an informative prior model in the latent space to further improve the expressive power of the whole model. This follows the philosophy of empirical Bayes where the prior model is learned from the observed data. Specifically, we assume the latent vector follows an energy-based model (EBM). We call this model the latent space energy-based prior model.
Both the latent space EBM and the top-down network can be learned jointly by maximum likelihood estimate (MLE). Each learning iteration involves Markov chain Monte Carlo (MCMC) sampling of the latent vector from both the prior and posterior distributions. Parameters of the prior model can then be updated based on the statistical difference between samples from the two distributions. Parameters of the top-down network can be updated based on the samples from the posterior distribution as well as the observed data.
Due to the low-dimensionality of the latent space, the energy function can be parametrized by a small multi-layer perceptron, yet the energy function can capture regularities in the data effectively because the EBM stands on an expressive top-down network. Moreover, MCMC in the latent space for both prior and posterior sampling is efficient and mixes well. Specifically, we employ short-run MCMC which runs a fixed number of steps from a fixed initial distribution. We formulate the resulting learning algorithm as a perturbation of MLE learning in terms of both objective function and estimating equation, so that the learning algorithm has a solid theoretical foundation. Within our theoretical framework, the short-run MCMC for posterior and prior sampling can also be amortized by jointly learned inference and synthesis networks. However, in this initial paper, we prefer keeping our model and learning method pure and self-contained, without mixing in learning tricks from variational auto-encoder (VAE) and generative adversarial networks (GAN) . Thus we shall rely on short-run MCMC for simplicity. The one-page code can be found in supplementary materials. See our follow-up development on amortized inference in the context of semi-supervised learning .
We test the proposed modeling, learning and computing method on tasks such as image synthesis, text generation, as well as anomaly detection. We show that our method is competitive with prior art. See also our follow-up work on molecule generation .
Contributions. (1) We propose a latent space energy-based prior model that stands on the top-down network of the generator model. (2) We develop the maximum likelihood learning algorithm that learns the EBM prior and the top-down network jointly based on MCMC sampling of the latent vector from the prior and posterior distributions. (3) We further develop an efficient modification of MLE learning based on short-run MCMC sampling. (4) We provide theoretical foundation for learning based on short-run MCMC. The theoretical formulation can also be used to amortize short-run MCMC by extra inference and synthesis networks. (5) We provide strong empirical results to illustrate the proposed method.
Model and learning
where is the prior model with parameters , is the top-down generation model with parameters , and .
The prior model is formulated as an energy-based model,
where is a known reference distribution, assumed to be isotropic Gaussian in this paper. is the negative energy and is parameterized by a small multi-layer perceptron with parameters . is the normalizing constant or partition function.
The prior model (2) can be interpreted as an energy-based correction or exponential tilting of the original prior distribution , which is the prior distribution in the generator model in VAE.
where , so that . As in VAE, takes an assumed value. For text modeling, let where each is a token. Following previous text VAE model , we define as a conditional autoregressive model,
which is parameterized by a recurrent network with parameters .
In the original generator model, the top-down network maps the unimodal prior distribution to be close to the usually highly multi-modal data distribution. The prior model in (2) refines so that maps the prior model to be closer to the data distribution. The prior model does not need to be highly multi-modal because of the expressiveness of .
The marginal distribution is The posterior distribution is
In the above model, we exponentially tilt . We can also exponentially tilt to . Equivalently, we may also exponentially tilt , as the mapping from to is a change of variable. This leads to an EBM in both the latent space and data space, which makes learning and sampling more complex. Therefore, we choose to only tilt and leave as a directed top-down generation model.
2 Maximum likelihood
Suppose we observe training examples . The log-likelihood function is
The learning gradient can be calculated according to
See Theoretical derivations in the Supplementary for a detailed derivation.
For the prior model, Thus the learning gradient for an example is
The above equation has an empirical Bayes nature. is based on the empirical observation , while is the prior model. is updated based on the difference between inferred from empirical observation , and sampled from the current prior.
where or for image and text modeling respectively.
Expectations in (7) and (8) require MCMC sampling of the prior model and the posterior distribution . We can use Langevin dynamics . For a target distribution , the dynamics iterates , where indexes the time step of the Langevin dynamics, is a small step size, and is the Gaussian white noise. can be either or . In either case, can be efficiently computed by back-propagation.
3 Short-run MCMC
Convergence of Langevin dynamics to the target distribution requires infinite steps with infinitesimal step size, which is impractical. We thus propose to use short-run MCMC for approximate sampling.
The short-run Langevin dynamics is always initialized from the fixed initial distribution , and only runs a fixed number of steps, e.g., ,
We then update and based on (10) and (11), where the expectations can be approximated by Monte Carlo samples. See our follow-up work on persistent chains for prior sampling.
4 Algorithm
The learning and sampling algorithm is described in Algorithm 1.
The posterior sampling and prior sampling correspond to the positive phase and negative phase of latent EBM .
5 Theoretical understanding
In terms of objective function, define the Kullback-Leibler divergence . At iteration , with fixed , consider the following computationally tractable perturbation of the log-likelihood function of for an observation ,
The above is a function of , while is fixed. Then
where is the total log-likelihood defined in equation (5), and the gradient is taken at .
In terms of estimating equation, the stochastic gradient descent in Algorithm 1 is a Robbins-Monro stochastic approximation algorithm that solves the following estimating equation:
6 Amortized inference and synthesis
In this initial paper, we prefer keeping our model and learning method clean and simple, without involving extra networks for learned computations, and without mixing in learning tricks from VAE and GAN. See our follow-up work on joint training of amortized inference network . See also for a temporal difference MCMC teaching scheme for amortizing MCMC.
Experiments
We present a set of experiments which highlight the effectiveness of our proposed model with (1) excellent synthesis for both visual and textual data outperforming state-of-the-art baselines, (2) high expressiveness of the learned prior model for both data modalities, and (3) strong performance in anomaly detection. For image data, we include SVHN , CelebA , and CIFAR-10 . For text data, we include PTB , Yahoo , and SNLI . We refer to the Supplementary for details. Code to reproduce the reported results is available https://bpucla.github.io/latent-space-ebm-prior-project/. Recently we extend our work to construct a symbol-vector coupling model for semi-supervised learning and learn it with amortized inference for posterior inference and persistent chains for prior sampling , which demonstrates promising results in multiple data domains. In another followup , we find the latent space EBM can learn to capture complex chemical laws automatically and implicitly, enabling valid, novel, and diverse molecule generations. Besides the results detailed below, our extended experiments also corroborate our modeling strategy of building a latent space EBM for powerful generative modeling, meaningful representation learning, and stable training.
We evaluate the quality of the generated and reconstructed images. If the model is well-learned, the latent space EBM will fit the generator posterior which in turn renders realistic generated samples as well as faithful reconstructions. We compare our model with VAE and SRI which assume a fixed Gaussian prior distribution for the latent vector and two recent strong VAE variants, 2sVAE and RAE , whose prior distributions are learned with posterior samples in a second stage. We also compare with multi-layer generator (i.e., 5 layers of latent vectors) model which admits a powerful learned prior on the bottom layer of latent vector. We follow the protocol as in .
Generation. The generator network in our framework is well-learned to generate samples that are realistic and share visual similarities as the training data. The qualitative results are shown in Figure 2. We further evaluate our model quantitatively by using Fréchet Inception Distance (FID) in Table 1. It can be seen that our model achieves superior generation performance compared to listed baseline models.
Reconstruction. We evaluate the accuracy of the posterior inference by testing image reconstruction. The well-formed posterior Langevin should not only help to learn the latent space EBM model but also match the true posterior of the generator model. We quantitatively compare reconstructions of test images with the above baseline models on mean square error (MSE). From Table 1, our proposed model could achieve not only high generation quality but also accurate reconstructions.
2 Text modeling
We compare our model to related baselines, SA-VAE , FB-VAE , and ARAE . SA-VAE optimized posterior samples with gradient descent guided by EBLO, resembling the short run dynamics in our model. FB-VAE is the SOTA VAE for text modeling. While SA-VAE and FB-VAE assume a fixed Gaussian prior, ARAE adversarially learns a latent sample generator as an implicit prior distribution. To evaluate the quality of the generated samples, we follow and recruit Forward Perplexity (FPPL) and Reverse Perplexity (RPPL). FPPL is the perplexity of the generated samples evaluated under a language model trained with real data and measures the fluency of the synthesized sentences. RPPL is the perplexity of real data computed under a language model trained with the model-generated samples. Prior work employs it to measure the distributional coverage of a learned model, in our case, since a model with a mode-collapsing issue results in a high RPPL. FPPL and RPPL are displayed in Table 2. Our model outperforms all the baselines on the two metrics, demonstrating the high fluency and diversity of the samples from our model. We also evaluate the reconstruction of our model against the baselines using negative log-likelihood (NLL). Our model has a similar performance as that of FB-VAE and ARAE, while they all outperform SA-VAE.
3 Analysis of latent space
We examine the exponential tilting of the reference prior through Langevin samples initialized from with target distribution . As the reference distribution is in the form of an isotropic Gaussian, we expect the energy-based correction to tilt into an irregular shape. In particular, learning equation 10 may form shallow local modes for . Therefore, the trajectory of a Markov chain initialized from the reference distribution with well-learned target should depict the transition towards synthesized examples of high quality while the energy fluctuates around some constant. Figure 3 and Table 3 depict such transitions for image and textual data, respectively, which are both based on models trained with steps. For image data the quality of synthesis improve significantly with increasing number of steps. For textual data, there is an enhancement in semantics and syntax along the chain, which is especially clear from step 0 to 40 (see Table 3).
While our learning algorithm recruits short run MCMC with steps to sample from target distribution , a well-learned should allow for Markov chains with realistic synthesis for steps. We demonstrate such long-run Markov chain with and in Figure 4. The long-run chain samples in the data space are reasonable and do not exhibit the oversaturating issue of the long-run chain samples of recent EBM in the data space (see oversaturing examples in Figure 3 in ).
4 Anomaly detection
We evaluate our model on anomaly detection. If the generator and EBM are well learned, then the posterior would form a discriminative latent space that has separated probability densities for normal and anomalous data. Samples from such a latent space can then be used to detect anomalies. We take samples from the posterior of the learned model, and use the unnormalized log-posterior as our decision function.
Following the protocol as in , we make each digit class an anomaly and consider the remaining 9 digits as normal examples. Our model is trained with only normal data and tested with both normal and anomalous data. We compare with the BiGAN-based anomaly detection , MEG and VAE using area under the precision-recall curve (AUPRC) as in . Table 4 shows the results.
5 Computational cost
Our method involving MCMC sampling is more costly than VAEs with amortized inference. Our model is approximately 4 times slower than VAEs on image datasets. On text datasets, ours does not have an disadvantage compared to VAEs on total training time (despite longer per-iteration time) because of better posterior samples from short run MCMC than amortized inference and the overhead of the techniques that VAEs take to address posterior collapse. To test our method’s scalability, we trained a larger generator on CelebA (). It produced faithful samples (see Figure 1).
Discussion and conclusion
We now put our work within the bigger picture of modeling and learning, and discuss related work.
Energy-based model and top-down generation model. A top-down model or a directed acyclic graphical model is of a simple factorized form that is capable of ancestral sampling. The prototype of such a model is factor analysis , which has been generalized to independent component analysis , sparse coding , non-negative matrix factorization , etc. An early example of a multi-layer top-down model is the generation model of Helmholtz machine . An EBM defines an unnormalized density or a Gibbs distribution. The prototypes of such a model are exponential family distribution, the Boltzmann machine , and the FRAME (Filters, Random field, And Maximum Entropy) model . contrasted these two classes of models, calling the top-down latent variable model the generative model, and the energy-based model the descriptive model. proposed to integrate the two models, where the top-down generation model generates textons, while the EBM prior accounts for the perceptual organization or Gestalt laws of textons. Our model follows such a plan. Recently, DVAEs adopted restricted Boltzmann machines as the prior model for binary latent variables and a deep neural network as the top-down generation model.
The energy-based model can be translated into a classifier and vice versa via the Bayes rule . The energy function in the EBM can be viewed as an objective function, a cost function, or a critic . It captures regularities, rules or constrains. It is easy to specify, although optimizing or sampling the energy function requires iterative computation such as MCMC. The maximum likelihood learning of EBM can be interpreted as an adversarial scheme , where the MCMC serves as a generator or an actor and the energy function serves as an evaluator or a critic. The top-down generation model can be viewed as an actor that directly generates samples. It is easy to sample from, though a complex top-down model is necessary for high quality samples. Comparing the two models, the scalar-valued energy function can be more expressive than the vector-valued top-down network of the same complexity, while the latter is much easier to sample from. It is thus desirable to let EBM take over the top layers of the top-down model to make it more expressive and make EBM learning feasible.
Energy-based correction of top-down model. The top-down model usually assumes independent nodes at the top layer and conditional independent nodes at subsequent layers. We can introduce energy terms at multiple layers to correct the independence or conditional independence assumptions, and to introduce inductive biases. This leads to a latent energy-based model. However, unlike undirected latent EBM, the energy-based correction is learned on top of a directed top-down model, and this can be easier than learning an undirected latent EBM from scratch. Our work is a simple example of this strategy where we correct the prior distribution. We can also correct the generation model in the data space.
From data space EBM to latent space EBM. EBM learned in data space such as image space can be highly multi-modal, and MCMC sampling can be difficult. We can introduce latent variables and learn an EBM in latent space, while also learning a mapping from the latent space to the data space. Our work follows such a strategy. Earlier papers on this strategy are . Learning EBM in latent space can be much more feasible than in data space in terms of MCMC sampling, and much of past work on EBM can be recast in the latent space.
Short-run MCMC and amortized computation. Recently, proposed to use short-run MCMC to sample from the EBM in data space. used it to sample the latent variables of a top-down generation model from their posterior distribution. used it to improve the posterior samples from an inference network. Our work adopts short-run MCMC to sample from both the prior and the posterior of the latent variables. We provide theoretical foundation for the learning algorithm with short-run MCMC sampling. Our theoretical formulation can also be used to jointly train networks that amortize the MCMC sampling from the posterior and prior distributions.
Generator model with flexible prior. The expressive power of the generator network for image and text generation comes from the top-down network that maps a simple prior to be close to the data distribution. Most of the existing papers assume that the latent vector follows a given simple prior, such as isotropic Gaussian distribution or uniform distribution. However, such assumption may cause ineffective generator learning as observed in . Some VAE variants attempted to address the mismatch between the prior and the aggregate posterior. VampPrior parameterized the prior based on the posterior inference model, while proposed to construct priors using rejection sampling. ARAE learned an implicit prior with adversarial training. Recently, some papers used two-stage approach . They first trained a VAE or deterministic auto-encoder. To enable generation from the model, they fitted a VAE or Gaussian mixture to the posterior samples from the first-stage model. VQ-VAE adopted a similar approach and an autoregressive distribution over was learned from the posterior samples. All of these prior models generally follow the empirical Bayes philosophy, which is also one motivation of our work.
2 Conclusion
EBM has many applications, however, its soundness and its power are limited by the difficulty with MCMC sampling. By moving from data space to latent space, and letting the EBM stand on an expressive top-down network, MCMC-based learning of EBM becomes sound and feasible, and EBM in latent space can capture regularities in data effectively. We may unleash the power of EBM in the latent space for many applications.
Broader Impact
Our work can be of interest to researchers working on generator model, energy-based models, MCMC sampling and unsupervised learning. It may also be of interest to people who are interested in image synthesis and text generation.
Acknowledgments and Disclosure of Funding
We thank the four reviewers for their insightful comments and useful suggestions. The work is supported by NSF DMS-2015577; DARPA XAI project N66001-17-2-4029; ARO project W911NF1810296; ONR MURI project N00014-16-1-2007; and XSEDE grant ASC170063. We thank the NVIDIA cooperation for the donation of 2 Titan V GPUs.
References
Appendix A Theoretical derivations
In this section, we shall derive most of the equations in the main text. We take a step by step approach, starting from simple identities or results, and gradually reaching the main results. Our derivations are unconventional, but they pertain more to our model and learning method.
Let . A useful identity is
where (or ) is the expectation with respect to .
The above identity has generalized versions, such as the one underlying the policy gradient , . By letting , we get (18).
A.2 Maximum likelihood estimating equation
The simple identity (18) also underlies the consistency of MLE. Suppose we observe independently, where is the true value of . The log-likelihood is
The maximum likelihood estimating equation is
According to the law of large number, as , the above estimating equation converges to
where is the unknown value to be solved, while is fixed. According to the simple identity (18), is the solution to the above estimating equation (22), no matter what is. Thus with regularity conditions, such as identifiability of the model, the MLE converges to in probability.
The optimality of the maximum likelihood estimating equation among all the asymptotically unbiased estimating equations can be established based on a further generalization of the simple identity (18).
We shall justify our learning method with short-run MCMC in terms of an estimating equation, which is a perturbation of the maximum likelihood estimating equation (21).
A.3 MLE learning gradient for θ𝜃\theta
Recall that , where . The learning gradient for an observation is as follows:
The above identity is a simple consequence of the simple identity (18).
because of the fact that according to the simple identity (18), while because what is inside the expectation only depends on , but does not depend on .
The above identity (23) is related to the EM algorithm , where is the observed data, is the missing data, and is the complete-data log-likelihood.
A.4 MLE learning gradient for α𝛼\alpha
For the prior model , we have . Applying the simple identity (18), we have
Hence the derivative of the log-likelihood is
According to equation (23) in the previous subsection, the learning gradient for is
We shall provide a theoretical understanding of the learning method with short-run MCMC in terms of Kullback-Leibler divergences. We start from some simple results.
The simple identity (18) also follows from Kullback-Leibler divergence. Consider
as a function of with fixed. Suppose the model is identifiable, then achieves its minimum 0 at , thus . Meanwhile,
Since is arbitrary in the above derivation, we can replace it by a generic , i.e.,
As a notational convention, for a function , we write , i.e., the derivative of at .
We now re-derive MLE learning gradient in terms of perturbation of log-likelihood by Kullback-Leibler divergence terms. Then the learning method with short-run MCMC can be easily understood.
At iteration , fixing , we want to calculate the gradient of the log-likelihood function for an observation , , at . Consider the following computationally tractable perturbation of the log-likelihood
In the above, as a function of , with fixed, is minimized at , thus its derivative at is 0. As a function of , with fixed, is minimized at , thus its derivative at is 0. Thus
We now unpack to see that it is computationally tractable, and we can obtain its derivative at .
where term gets canceled,
do not depend on . consists of two entropy terms. Now taking derivative at , we have
Averaging over the observed examples leads to MLE learning gradient.
In the above, we calculate the gradient of at . Since is arbitrary in the above derivation, if we replace by a generic , we get the gradient of at a generic , i.e.,
The above calculations are related to the EM algorithm and the learning of energy-based model.
In EM algorithm, the complete-data log-likelihood serves as a surrogate for the observed-data log-likelihood , where
and , where is a lower-bound of or minorizes the latter. and touch each other at , and they are co-tangent at . Thus the derivative of at is the same as the derivative of at .
In EBM, serves to cancel term in the EBM prior, and is related to the second divergence term in contrastive divergence .
A.7 Maximum likelihood estimating equation for θ=(α,β)𝜃𝛼𝛽\theta=(\alpha,\beta)
Based on (48) and (49), the estimating equation is
A.8 Learning with short-run MCMC as perturbation of log-likelihood
At iteration , fixing , the updating rule based on short-run MCMC follows the gradient of the following function, which is a perturbation of log-likelihood for the observation ,
The above is a function of , while is fixed.
In full parallel to the above subsection, we have
Averaging over , we get the updating rule based on short-run MCMC. That is, the learning rule based on short-run MCMC follows the gradient of a perturbation of the log-likelihood function where the perturbations consists of two terms.
A.9 Perturbation of maximum likelihood estimating equation
The fixed point of the learning algorithm based on short-run MCMC is where the update is 0, i.e.,
This is clearly a perturbation of the MLE estimating equation in (52) and (53). The above estimating equation defines an estimator, where the learning algorithm with short-run MCMC converges.
We can rewrite the objective function (54) in a more revealing form. Let independently, where is the data distribution. At time step , with fixed , learning based on short-run MCMC follows the gradient of
Let us assume is large enough, so that the average is practically the expectation with respect to . Then MLE maximizes , which is equivalent to minimizing . The learning with short-run MCMC follows the gradient that minimizes
where, with some abuse of notation, we now define
where we also average over , instead fixing as before.
where the on the right hand side is about the joint distributions of , and is more tractable than the first on the left hand side, which is for MLE. This underlies EM and VAE. Now subtracting the third , we have the following special form of contrastive divergence
As mentioned in the main text, we can also exponentially tilt to , or equivalently, exponentially tilt . The above derivations can be easily adapted to such a model, which we choose not to explore due to the complexity of EBM in the data space.
A.11 Amortized inference and synthesis networks
We can then define the following objective function in parallel with the objective function (61) in the above subsection,
and we can jointly learn , and by
See for related formulations. The learning of the inference network follows VAE. The learning of the synthesis network is based on variational approximation to . The pair and play adversarial roles, where serves as an actor and serves as a critic.
Appendix B Experiments
Data. Image datasets include SVHN , CIFAR-10 , and CelebA . We use the full training split of SVHN () and CIFAR-10 () and take examples of CelebA as training data following . The training images are resized and scaled to $$. Text datasets include PTB , Yahoo , and SNLI , following recent work on text generative modeling with latent variables .
Model architectures. The architecture of the EBM, , is displayed in Table 6. For text data, the dimensionality of is set to . The generator architectures for the image data are also shown in Table 6. The generators for the text data are implemented with a one-layer unidirectional LSTM and Table 7 lists the number of word embeddings and hidden units of the generators for each dataset.
Short run dynamics. The hyperparameters for the short run dynamics are depicted in Table 5 where and denote the number of prior and posterior sampling steps with step sizes and , respectively. These are identical across models and data modalities, except for the model for CIFAR-10 which is using steps.
Optimization. The parameters for the EBM and image generators are initialized with Xavier normal and those for the text generators are initialized from a uniform distribution, Unif, following . Adam is adopted for all model optimization. The models are trained until convergence (taking approximately and parameter updates for image and text models, respectively).
Appendix C Ablation study
We investigate a range of factors that are potentially affecting the model performance with SVHN as an example. The highlighted number in Tables 8, 9, and 10 is the FID score reported in the main text and compared to other baseline models. It is obtained from the model with the architecture and hyperparameters specified in Table 5 and Table 6 which serve as the reference configuration for the ablation study.
Fixed prior. We examine the expressivity endowed with the EBM prior by comparing it to models with a fixed isotropic Gaussian prior. The results are displayed in Table 8. The model with an EBM prior clearly outperforms the model with a fixed Gaussian prior and the same generator as the reference model. The fixed Gaussian models exhibit an enhancement in performance as the generator complexity increases. They however still have an inferior performance compared to the model with an EBM prior even when the fixed Gaussian prior model has a generator with four times more parameters than that of the reference model.
MCMC steps. We also study how the number of short run MCMC steps for prior inference () and posterior inference (). The left panel of Table 9 shows the results for and the right panel for . As the number of MCMC steps increases, we observe improved quality of synthesis in terms of FID.
Prior EBM and generator complexity. Table 10 displays the FID scores as a function of the number of hidden features of the prior EBM (nef) and the factor of the number of channels of the generator (ngf, also see Table 6). In general, enhanced model complexity leads to improved generation.