Inference with Deep Generative Priors in High Dimensions
Parthe Pandit, Mojtaba Sahraee-Ardakan, Sundeep Rangan, Philip Schniter, Alyson K. Fletcher
I Introduction
We consider inference in an -layer stochastic neural network of the form
The inference problem (2) arises in the following state-of-the-art approach to inverse problems. In general, solving an “inverse problem" means recovering some signal from a measurement that depends on . For example, in compressed sensing (CS) , the measurements are often modeled as with known and additive white Gaussian noise (AWGN) , and the signal is often modeled as a sparse linear combination of elements from a known dictionary, i.e., for some sparse coefficient vector . To recover , one usually computes a sparse coefficient estimate using a LASSO-type convex optimization and then uses it to form a signal estimate , as in
where is a tunable parameter. The CS recovery approach (3) can be interpreted as a two-layer version of the inference problem: the first layer implements signal generation via , while the second layer implements the measurement process . Equation (3) then performs maximum a posteriori inference (see the discussion around (6)) to recover estimates of and .
Although CS has met with some success, it has a limited ability to exploit the complex structure of natural signals, such as images, audio, and video. This is because the model “ with sparse ” is overly simplistic; it is a one-layer generative model. Much more sophisticated modeling is possible with multi-layer priors, as demonstrated in recent works on variational autoencoders (VAEs) , generative adversarial networks (GANs) , and deep image priors (DIP) . These models have had tremendous success in modeling richly structured data, such as images and text.
A typical application of solving an inverse problem using a deep generative model is shown in Fig. 2. This figure considers the classic problem of inpainting , for which reconstruction with DIP has been particularly successful . Here, a noise-like signal drives a three-layer generative network to produce an image . The generative network would have been trained on an ensemble of images similar to the one being estimated using, e.g., VAE or GAN techniques. The measurement process, which manifests as occlusion in the inpainting problem, is modeled using one additional layer of the network, which produces the measurement . Inference is then used to recover the image (i.e., the hidden-layer signal ) from . In addition to inpainting, this deep-reconstruction approach can be applied to other linear inverse problems (e.g., CS, de-blurring, and super-resolution) as well as generalized-linear inverse problems (e.g., classification, phase retrieval, and estimation from quantized outputs). We note that the inference approach provides an alternative to designing and training a separate reconstruction network, such as in .
When using deterministic deep generative models, the unknown signal can be modeled as , where is a trained deep neural network and is a realization of an i.i.d. random vector, typically with a Gaussian distribution. Consequently, to recover from a linear-AWGN measurement of the form , the compressed-sensing approach in (3) can be extended to a regularized least-squares problem of the form
In practice, the optimization in (4) is solved using a gradient-based method. This approach can be straightforwardly implemented with deep-learning software packages and has been used, with excellent results, in . The minimization (4) has also been useful in interpreting the semantic meaning of hidden signals in deep networks . VAEs and certain GANs can also produce decoding networks that sample from the posterior density, and sampling methods such as Markov-chain Monte Carlo (MCMC) algorithms and Langevin diffusion can also be employed.
I-B Analysis via Approximate Message Passing (AMP)
While reconstruction with deep generative priors has seen tremendous practical success, its performance is not fully understood. Optimization approaches such as (4) are typically non-convex and difficult to analyze. As we discuss below, most results available today only provide bounds, and these bounds are often be overly conservative (see Section I-D).
To answer these questions, this paper considers deep inference via approximate message passing (AMP), a powerful approach for analyzing estimation problems in certain high-dimensional random settings. Since its origins in understanding linear inverse problems in compressed sensing , AMP has been extended to an impressive range of estimation and learning tasks, including generalized linear models , models with parametric uncertainty , structured priors, and bilinear problems. For these problems, AMP-based methods have been able to provide computationally efficient algorithms with precise high-dimensional analyses. Often, AMP approaches yield optimality guarantees in cases where all other known approaches do not.
I-C Main Contributions
We establish several key results on the ML-VAMP algorithm:
We show that, for both MAP and MMSE inference, the fixed points of the ML-VAMP algorithm correspond to stationary points of variational formulations of these estimators. This allows the interpretation of ML-VAMP as a Lagrangian algorithm with adaptive step-sizes in both cases. These findings are given in Theorems 1 and 2 and are similar to previous results for AMP . Section III describes these results.
We prove that, in a certain large system limit (LSL), the behavior of ML-VAMP is exactly described by a deterministic recursion called the state evolution (SE). This SE analysis is a multi-layer extension of similar results for AMP and VAMP. The SE equations enable asymptotically exact predictions of macroscopic behaviors of the hidden-layer estimates for each iteration of the ML-VAMP algorithm. This allows us to obtain error bounds even if the algorithm is run for a finite number of iterations. The SE analysis, given in Theorem 3, is the main contribution of the paper, and is discussed in Section IV.
Since the original conference versions of this paper , formulae for the minimum mean-squared error (MMSE) for inference in deep networks have been conjectured in . As discussed in Section IV-C, these formulae are based on heuristic techniques, such as the replica method from statistical physics, and have been rigorously proven in special cases . Remarkably, we show that the mean-squared-error (MSE) of ML-VAMP exactly matches the predicted MMSE in certain cases.
Using numerical simulations, we verify the predictions of the main result from Theorem 3. In particular, we show that the SE accurately predicts the MSE even for networks that are not considered large by today’s standards. We also perform experiments with the MNIST handwritten digit dataset. Here we consider the inference problem using learned networks, for which the weights do not satisfy the randomness assumptions required in our analysis.
In summary, ML-VAMP provides a computationally efficient method for inference in deep networks whose performance can be exactly predicted in certain high-dimensional random settings. Moreover, in these settings, the MSE performance of ML-VAMP can match the existing predictions of the MMSE.
I-D Prior Work
There has been growing interest in studying learning and inference problems in high-dimensional, random settings. One common model is the so-called wide network, where the dimensions of the input, hidden layers, and output are assumed to grow with a fixed linear scaling, and the weight matrices are modeled as realizations of random matrices. This viewpoint has been taken in , in several works that explicitly use AMP methods , and in several works that use closely related random-matrix techniques .
The existing work most closely related to ours is that by Manoel et al. , which developed a multi-layer version of the original AMP algorithm . The work provides a state-evolution analysis of multi-layer inference in networks with entrywise i.i.d. Gaussian weight matrices. In contrast, our results apply to the larger class of rotationally invariant matrices (see Section IV for details), which includes i.i.d. Gaussian matrices case as a special case.
Some of the material in this paper appeared in conference versions . The current paper includes all the proofs, simulation details, and provides a unified treatment of both MAP and MMSE estimation.
II Multi-layer Vector Approximate Message Passing
Similar to other graphical-model methods , we consider two forms of estimation: MAP estimation and MMSE estimation. The maximum a priori, or MAP, estimate is defined as
Although we will focus on MAP estimation, most of our results will apply to general -estimators of the form,
We will also consider the minimum mean-squared error, or MMSE, estimate, defined as
II-B The ML-VAMP Algorithm
II-C MAP and MMSE Estimation Functions
Given these parameters, both the MAP and MMSE estimation functions are defined from the belief function
II-D Computational Complexity
From Algorithm 1, we see that each pass of the MAP-ML-VAMP or MMSE-ML-VAMP algorithm requires solving (a) scalar MAP or MMSE estimation problems for the non-linear, separable layers; and (b) least-squares problems for the linear layers. In particular, no high-dimensional integrals or high-dimensional optimizations are involved.
III Fixed Points of ML-VAMP
Applying these relationships to lines 10 and 20 of Algorithm 1 gives
Corresponding to this constrained optimization, we define the augmented Lagrangian
whereas the backward pass iterations satisfy
Further, any fixed point of Algorithm 1 corresponds to a critical point of the Lagrangian (20).
III-B Fixed Points of MMSE-ML-VAMP and Connections to Free-Energy Minimization
Because , we can express it using variational optimization as
yields a tractable approximation to .
IV Analysis in the Large-System Limit
ML-VAMP algorithm
Distribution of the components
State Evolution
IV-B SE Analysis in the LSL
Under these assumptions, we can now state our main result.
IV-C MMSE Estimation and Connections to the Replica Predictions
Since the estimation functions in Theorem 4 are the MSE optimal functions for true densities, we will call this selection of estimation functions the MMSE matched estimators. Under the assumption of MMSE matched estimators, the theorem shows that the MSE error has a simple set of recursive expressions.
It is useful to compare the predicted MSE with the predicted optimal values. The works postulate the optimal MSE for inference in deep networks under the LSL model described above using the replica method from statistical physics. Interestingly, it is shown in [47, Thm.2] that the predicted minimum MSE satisfies equations that exactly agree with the fixed points of the updates (39). Thus, when the fixed points of (39) are unique, ML-VAMP with matched MMSE estimators provably achieves the Bayes optimal MSE predicted by the replica method. Although the replica method is not rigorous, this MSE predictions have been indepedently proven for the Gaussian case in and certain two layer networks in . This situation is similar to several other works relating the MSE of AMP with replica predictions . The consequence is that, if the replica method is correct, ML-VAMP provides a computationally efficient method for inference with testable conditions under which it achieves the Bayes optimal MSE.
V Numerical Simulations
We now numerically investigate the MAP-ML-VAMP and MMSE-ML-VAMP algorithms using two sets of experiments, where in each case the goal was to solve an estimation problem of the form in (2) using a neural network of the form in (1). We used the Python 3.7 implementation of the ML-VAMP algorithm available on GitHub.See https://github.com/GAMPTeam/vampyre.
To quantify the performance of ML-VAMP, we repeated the following 1000 times. First, we drew a random neural network as described above. Then we ran the ML-VAMP algorithm for 100 iterations, recording the normalized MSE (in dB) of the iteration- estimate of the network input, :
Since ML-VAMP computes two estimates of at each iteration, we consider each estimate as corresponding to a “half iteration.”
For MMSE-ML-VAMP, the left panel of Fig. 3 shows the NMSE versus half-iteration for 100 compressed measurements. The value shown is the average over 1000 random realizations. Also shown is the MSE predicted by the ML-VAMP state evolution. Comparing the two traces, we see that the SE predicts the actual behavior of MMSE-ML-VAMP remarkably well, within approximately 1 dB. The right panel shows the NMSE after 50 iterations (i.e., 100 half-iterations) for several numbers of measurements . Again we see an excellent agreement between the actual MSE and the SE prediction. In both cases we used the positive fraction . Analogous results are shown for MAP-ML-VAMP in Fig. 4. There we see an excellent agreement between the actual MSE and the SE prediction for iterations and all values of .
Comparison to ADAM
We now compare the MSE of MAP-ML-VAMP and its SE to that the MAP approach (4) using the ADAM optimizer , as implemented in Tensorflow. As before, the goal was to recover the input to the 7-layer synthetic network from a measurement of its output. Fig. 5 shows the median NMSE over 40 random network realizations for several values of , the number of measurements. We see that, for , the performance of MAP-ML-VAMP closely matches its SE prediction, as well as the performance of the ADAM-based MAP approach (4). For , there is a discrepancy between the MSE performance of MAP-ML-VAMP and its SE prediction, which is likely due to the relatively small dimensions involved. Also, for small , MAP-ML-VAMP appears to achieve slightly better MSE performance than the ADAMP-based MAP approach (4). Since both are attempting to solve the same problem, the difference is likely due to ML-VAMP finding better local minima.
V-B Image Inpainting: MNIST dataset
To demonstrate that ML-VAMP can also work on a real-world dataset, we perform inpainting on the MNIST dataset. The MNIST dataset consists of 28 28 = 784 pixel images of handwritten digits, as shown in the first column of Fig. 6.
To start, we trained a 4-layer (deterministic) deep generative prior model from 50 000 digits using a variational autoencoder (VAE) . The VAE “decoder” network was designed to accept 20-dimensional i.i.d. Gaussian random inputs with zero mean and unit variance, and to produce MNIST-like images . In particular, this network began with a linear layer with 400 outputs, followed by a ReLU activations, followed by a linear layer with 784 units, followed by sigmoid activations that forced the final pixel values to between 0 and 1.
Given an image, , our measurement process produced by erasing rows 10-20 of , as shown in the second column of Fig. 6. This process is known as “occlusion.” By appending the occlusion layer onto our deep generative prior, we got a 5-layer network that generates an occluded MNIST image from a random input . The “inpainting problem” is to recover the image from the occluded image .
VI Conclusion
Inference using deep generative prior models provides a powerful tool for complex inverse problems. Rigorous theoretical analysis of these methods has been difficult due to the non-convex nature of the models. The ML-VAMP methodology for MMSE as well as MAP estimation provides a principled and computationally tractable method for performing the inference whose performance can be rigorously and precisely characterized in a certain large system limit. The approach thus offers a new and potentially powerful approach for understanding and improving deep neural network based models for inference.
Appendix A Empirical Convergence of Vector Sequences
That is, applies the function on each -dimensional component. Similarly, we say acts componentwise on whenever it is of the form (40) for some function .
Next consider a sequence of block vectors of growing dimension,
where, as usual, we have omitted the dependence on in . Importantly, empirical convergence can be defined on deterministic vector sequences, with no need for a probability space. If is a random vector sequence, we will often require that the limit (42) holds almost surely.
Appendix B ML-VAMP State Evolution Equations
The state evolution (SE) recursively defines a set of scalar random variables that describe the typical components of the vector quantities produced from the ML-VAMP algorithm. The definition of the random variables are given in Algorithm 2. The algorithm steps mimic those in the ML-VAMP algorithm, Algorithm 1, but with each update producing scalar random variables instead of vectors. The updates use several functions:
Appendix C Proofs of ML-VAMP Fixed-Point Theorems
Indeed the above conditions are the stationarity conditions of the optimization problem in (22a) and (23a). Hence (47) holds.
C-B Proof of Theorem 2
Observe that the Lagrangian function for the constrained optimization problem (29) for this specific choice of Lagrange multipliers is given by
Appendix D General Multi-Layer Recursions
For vectors in the Gen-ML Algorithm (Algorithm 3), we assume:
to some list . The limit (52) means that every element in the list converges to a limit as almost surely.
We also assume the following about the behaviour of component functions around the quantities defined in Algorithm 4. The iteration index has been dropped for simplifying notation.
For component functions and parameter update functions we assume:
We are now ready to state the general result regarding the empirical convergence of the true and iterated vectors from Algorithm 3 in terms of random variables defined in Algorithm 4.
Consider the iterates of the Gen-ML recursion (Algorithm 3) and the corresponding random variables and parameter limits defined by the SE recursions (Algorithm 4) under Assumptions 1 and 2. Then,
Appendix E in the supplementary materials is dedicated to proving this result.
Appendix E Proof of Theorem 5
The proof is similar to that of [35, Theorem 4], which provides a SE analysis for VAMP on a single-layer network. The critical challenge here is to extend that proof to multi-layer recursions. Many of the ideas in the two proofs are similar, so we highlight only the key differences between the two.
We prove these hypotheses by induction via a sequence of implications,
The inductive step then corresponds to the following result.
(a) follows from Stein’s Lemma; and (b) follows from (68), and (69). Consequently,
The proof is very similar to that of [35, Lemmas 7,8].
Note that the above PL(2) convergence can be shown using the same arguments involved in showing that if and then for some constant and sigma-algebra .
Appendix F Proofs of Main Results: Theorems 3 and 4
Recall that the main result in Theorem 3 claims the empirical convergence of PL(2) statistics of iterates of the ML-VAMP algorithm 1 to the expectations corresponding statistics of random variables given in Algorithm 2. We prove this result by applying the general convergence result stated in Theorem 5 which shows that under Assumptions 1 and 2, the PL(2) statistics of iterates of Algorithm 3 empirically converge to expectations of corresponding statistics of appropriately defined scalar random variables defined in Algorithm 4.
The proof of Theorem 3 proceeds in two steps. First, we show that the ML-VAMP iterations are a special case of the iterations of Algorithm 3, and similarly Algorithm 2 is a special case of 2, for specific choices of vector update functions, parameter statistic functions and parameter update functions, and their componentwise counterparts. The second step is to show that all assumptions required in Theorem 5 are satisfied, and hence the conclusions of Theorem 5 hold.
We start by showing that the ML-VAMP iterations from Algorithm 1 are a special case of the Gen-ML recursions from Algorithm 3.
Algorithm 4 is equivalent to Algorithm 2.
Assumptions 1 and 2 are satisfied by the conditions in Theorem 3.
Theorem 5 leads to the conclusion that the following triplets are asymptotically normal
The results in Theorem 3 follows from the argument definition of PL(2) convergence defined in Appendix A