Recurrent Generative Adversarial Networks for Proximal Learning and Automated Compressive Image Recovery

Morteza Mardani, Hatef Monajemi, Vardan Papyan, Shreyas Vasanawala, David Donoho, John Pauly

Introduction

Linear inverse problems widely appear in image restoration tasks in applications ranging from super-resolving natural images to reconstructing biomedical images. In such applications, one oftentimes encounters a seriously ill-posed recovery task, which necessitates regularization with proper statistical priors. This is however impeded by the following challenges: c1) real-time and interactive tasks afford only a low overhead for inference and training; e.g., imagine MRI visualization for neurosurgery , or, real-time superresolution that may need re-training on a cell phone ; c2) the need for recovering plausible images that are consistent with the physical model; this is particularly important for medical diagnosis, which is sensitive to hallucination.

For biomedical images one typically knows a small fraction of projections onto a certain transform domain (e.g., Fourier, or, Radon) based on physics of the scanner. Variations of deep CNNs are trained to map out aliased MR (AutoMap ) or low-dose CT (RED-CNN ) images to the gold-standard ones retrieved by iterative CS. They offer rapid reconstruction at the expense of high training overhead. There is however no systematic mechanism to assure fidelity to the underlying physical model, which can possibly hallucinate images. Another line of work pertains to developing effective priors to incorporate in an iterative algorithm which can outperform the conventional sparsity priors for CS; see e.g., . For instance, uses the low-dimensional code offered by a pre-trained generative decoder to achieve higher SNR than CS. These schemes attain a reasonably high SNR, but need several iterations for convergence that hinders real-time imaging.

Toward rapid, feasible, and plausible image recovery for ill-posed linear inverse tasks, this paper proposes a novel approach to automate imaging. Inspired by proximal gradient iterations, a recurrent ResNet architecture is designed to learn the proximal(s) from the data. Image prior information is learned through the proximal, that is modeled as a generator (G) network, consisting a few residual blocks (RB), trained with mixture of pixel-wise and perceptual costs via GANs to recover plausible images. The overall architecture implements multiple back-and-forth (approximate) projections onto the subspace of data-consistent images - dictated by the physical model - and the manifold of plausible images. The number of projections and similarly the size of each G network is desired to be small to reduce the training and inference overhead for real-time and interactive image recovery tasks. To study this effect, we perform experiments for reconstructing pediatric MR images, and super-resolving natural images from the CelebA face dataset. The former is a global task with subsampled measurements in the frequency domain, so-termed kk-space, while the latter is a rather local task, with a pixelated low-resolution image.

Our observations indicate that for MRI reconstruction, one better repeat a small ResNet (with a single RB) multiple times, instead of training a deep network, as commonly used e.g., in . The recurrent architecture not only improves upon the deep schemes by about 22dB SNR, but also incurs much less training overhead. This simple architecture also significantly outperforms the conventional CS-Wavelet/CS-TV schemes; 44dB SNR gain in 100100 times shorter time. For the local 44x superresolution task however it turns out that one needs to train a deep ResNets to learn the denoising proximals, and the preliminary results do not show any major advantage in using recurrent schemes. Before closing this section, it is worth to re-iterate the main contributions of this paper as follows:

Learning proximals using ResNet generators trained with pixel-wise and GAN-based perceptual costs

Extensive experiments for MRI reconstruction and face superresolution gives insight to design proper network architectures for various recovery tasks

The rest of this paper is organized as follows. Section 2 states the problem. Proximal learning based on recurrent GANs is discussed in Section 3. Evaluations for pediatric MR image reconstruction, and natural image superresolution are reported in Section 4, while the conclusions are drawn in Section 5.

Preliminaries and problem statement

The stated problem covers a wide range of image restoration and reconstruction tasks. For instance, in medical image reconstruction Φ\bm{\Phi} describes a projection driven by physics of the acquisition system (e.g., Fourier transform for MRI scanner). For image superresolution it is the downsampling operator that averages out nonoverlapping image regions to arrive at a low-resolution one. Given the image prior distribution, one typically forms a maximum-likelihood estimator formulated as a regularized least-squares (LS) program

with the regularizer ψ(⋅)\psi(\cdot) parameterized by Θ\bm{\Theta} that incorporates the image prior.

In order to solve (P1) one can adopt a variation of proximal gradient algorithm with a proximal operator Pψ{⋅}\mathcal{P}_{\psi}\{\cdot\} that is obtained based on ψ(⋅)\psi(\cdot) . Starting from x=0\mathbf{x}=\mathbf{0}, and adopting a small step size α\alpha the overall iterative procedure is expressed as

As argued earlier FISTA is a universal and thus naive regularization that does not take into account the image complications and perceptual quality. In addition, for moderate and high resolution images it demands many iterations for convergence that can seriously imped real-time recovery. The next sections aim to fix these caveats by learning proximals from historical images using generative neural networks.

Proximal learning

Motivated by the proximal gradient iterations in (2), to design efficient network architectures that automatically invert linear tasks, we need to first address the following important questions:

How can one ensure the network does not hallucinate images, and retrieves plausible images that are physically feasible?

How can one ensure rapid inference and affordable training for real-time and interactive image recovery tasks?

To bypass this hurdle, inspired by recurrent neural networks (RNNs) we unroll the loop and repeat multiple, say KK, copies of the proximal network as depicted in Fig. 1 (bottom). Each proximal network is accompanied with an (approximate) data consistency projection that refines xˇ\check{\mathbf{x}} to be consistent with the observations by simply moving along the descent direction of the data fidelity cost, namely ∥y−Φx∥2\|\mathbf{y}-\bm{\Phi}\mathbf{x}\|^{2}. Assuming exact data consistency projection the unrolled network learns the projection onto the intersection of physically feasible and visually plausible images. In general one can consider the cascaded architecture in Fig. 1 with independent weights {Θk}k=1K\{\bm{\Theta}_{k}\}_{k=1}^{K} per copies, but we are more interested in sharing the weights, namely Θ1=…=ΘK\bm{\Theta}_{1}=\ldots=\bm{\Theta}_{K}, which needs less training variables, and the back-propagation can easily accommodate gradient calculations .

In essence, (multiple) back-and-forth projections can ensure data fidelity to a good extent. This is in contrast with the existing deep architectures for automated medical image reconstruction (e.g., ) with no consideration for data fidelity, which may hallucinate images and mislead the diagnosis. It is also worth mentioning that the Amortised-MAP based deep-GAN scheme in uses an affine projection layer that improves GAN’s stability and suerresolution quality. However, one naturally need multiple projections to assure data consistency. The number of copies however cannot be large, or, alternatively the proximal networks need to be small, for real-time inference tasks.

2 Mixture of pixel-wise and perceptual costs

To learn proximals as projections onto manifold of visually plausible images we adopt GANs . Conventional generative models such as variational auto-encoders rely on pixel-wise costs that offer high pick signal-to-noise ratios but often produce overly-smooth images with poor perceptual quality. GANs however train a perceptual loss from the training data. Standard GANs consist of a tandem structure of generator (G) and discriminator (D) networks .

Training GANs amounts to playing a game with conflicting objectives between the adversary G and the discriminator. D network aims to score one the training ground-truth images drawn from the data distribution, and zero the (fake) outputs of G. Apparently, D cannot perfectly separate real and fake images as G tries to generate fake images that fools G. Various strategies have been devised to reach the game’s equilibrium. They mostly differ in evaluating the loss incurred by G and D , . The conventional GAN uses a sigmoid cross-entropy for D’s loss, which suffers from vanishing gradients. It leads to unstable training that causes mode collapse. In addition, for the generated images classified confidently as real (with a large decision variable), no cost is incurred. Hence, it tends to pull samples away from the decision boundary, that introduces non-realistic images . This particularly can hallucinate medical images, and as a result mislead medical diagnosis. To alleviate this issue, we adopt least-square GAN (LSGAN) that penalizes the classification mistake with a LS cost that pulls the generated samples towards the decision boundary.

where ∥x∥1,2:=γ∥x∥1+(1−γ)∥x∥2\|\mathbf{x}\|_{1,2}:=\gamma\|\mathbf{x}\|_{1}+(1-\gamma)\|\mathbf{x}\|_{2} for some 0≤γ≤10\leq\gamma\leq 1. The LS data fidelity term in (P1.2) is a soft version of the affine projection in the network architecture of Fig. 1. Parameter λ\lambda is also tuned based on the measurement noise level and the expected pixel-wise fidelity.

Experiments

Performance of the novel recurrent GANCS scheme is assessed in reconstructing pediatric MR images and super-resolving natural images. The former introduces aliasing artifacts that globally impact the entire image pixels, while in the latter the pixelation occurs locally. While the focus is mostly placed on MRI, preliminary results are also reported for image super-resolution to shed some light on challenges associated with proximal learning. In particular, we aim to address the following intriguing questions:

Q1. What is the proper number of copies, and generator size to learn the proximal?

Q2. What is the trade-off between PSNR/SSIM and inference/training complexity?

Q3. How is the performance compared with the conventional sparse coding?

Q4. How does the performance change if we train with independent weights per copies, and what is the interpretation for output of different copies?

To address the above questions, for the generator networks we adopt a ResNet with a variable number of residual blocks (RB). Each RB consists of two convolutional layers with 3×33\times 3 kernels and a fixed number of 128128 feature maps, respectively, that are followed by batch normalization (BN) and ReLU activation. It is then followed by three simple convolutional layers with 1×11\times 1 kernels, where the first two layers undergo ReLU activation and the last layer has sigmoid activation to return the output; see Fig. 2. Notice that for all generators {G(Θk)}k=1K\{\mathcal{G}(\bm{\Theta}_{k})\}_{k=1}^{K} a similar ResNet architecture is used.

The D network is composed of eight convolutional layers. In all the layers except the last one, the convolution is followed by BN and ReLU activation. No pooling is used. For the first four layers, number of feature maps is doubled from 88 to 6464, while at the same time convolution with stride 22 is used to reduce the image resolution. Kernel size 3×33\times 3 is adopted for the first five layers, while the last two layers use kernel size 1×11\times 1. In the last layer, the convolution output is averaged out to form the decision variable for LS binary classification, where no soft-max is used.

Adam optimizer is used with the momentum parameter β=0.9\beta=0.9, mini-batch size Lb=2L_{b}=2, and learning rate μ=10−5\mu=10^{-5}. Training is performed with TensorFlow interface on a NVIDIA Titan X Pascal GPU with 12GB RAM.

2 MRI reconstruction and artifact suppression

Performance of the novel recurrent scheme is assessed in removing aliasing artifacts from MR images. In essence, the scanner acquires Fourier coefficients (kk-space data) of the underlying image across various coils. A single-coil MR acquisition model is considered where for nn-th patient the acquired kk-space data admits

Here, F\mathcal{F} refers to the 2D Fourier transform, and the set Ω\Omega indexes the sampled Fourier coefficients. As it is conventionally performed with CS MRI, we select Ω\Omega based on a variable density sampling with radial view ordering that is more likely to pick low frequency components from the center of kk-space . Only 20%20\% of Fourier coefficients are collected. The sampling mask is shown in Fig. 4.

In order to assess the impact of network wiring on the image recovery performance, the cascaded network is trained for a variable number of ResNet copies with variable number of RBs. 1010k slices from the train dataset set are randomly picked for training, and 1,2801,280 slices from the test dataset for test.

Independent weights. We also consider a scenario where one allows weights varying across different copies. A similar ResNet architecture is used for all copies, which multiplies the variable count for training by the number of copies. As seen in Fig. 6, adopting a single RB per copy and repeating it for 15−2015-20 copies seems to be a suitable choice that achieves up to 27.427.4dB SNR. Apparently, a single RB and 1010 copies performs as good as 4−54-5 RBs with 55 copies in terms of SNR and SSIM. Comparing with the shared weight scenario, for 1010 copies with a single RB, using independent weights improves the SNR by almost 11dB. Notice that when each copy includes more than 55 RBs, our GPU resources become exhausted for more than 55 copies, and thus the rest of points are not shown on the plot. Further evaluations with more efficient implementation and stronger GPU resources is deferred for our future research.

Training and inference time. Inference time for both shared and independent weights is the same, and proportional to the number of copies. Feed-forwarding each image through a copy with one RB takes 44 msec when fully using the GPU. The training variable count is also proportional to the number of copies when the weights are allowed to change per different copies. It is hard to precisely evaluate the training and inference time under fair conditions as it strongly depends on the implementation and the allocated memory and processing power per run. As an estimate for the inference time we average it out over a few runs on the GPU as listed in Table 2. It is empirically observed that with shared weights, e.g., 1010 copies with 11 RB the training converges in 2−32-3 hours, but a deep single copy ResNet with 1010 RBs takes around 10−1210-12 hours to converge.

2.2 Comparison with sparse coding

To compare with the conventional CS schemes, CS-WV and CS-TV are adopted and tunned for the best SNR performance using BART that runs 300300 iterations of FISTA along with 100100 iterations of conjugate gradient descent to reach convergence. Quantitative results are listed under Table 2, where it is evident that the recurrent scheme with shared weights significantly outperforms CS with more than 44dB SNR gain that leads to sharper images with finer texture details as seen in Fig. 8. As a representative example Fig. 8 depicts the reconstructed abdominal slice of a test patient. CS-WV retrieves a blurry image that misses out the sharp details of the liver vessels. A deep ResNet with one copy and 1010 RBs captures a cleaner image, but still smoothens out fine texture details such as vessels. However, when using 1010 simple copies with a single RB, more details are seen about the liver vessels, and the texture appears to be more realistic. Similarly, using 55 copies each containing 22 RBs retrieves finer details than 22 relatively large copies with 55 RBs.

This observation indicates that the proximal for denoising MR images is well represented by a small number 1−21-2 RBs. The important message however is that multiple back-and-forth iterations are needed to recover a plausible MR image that is physically feasible. considering the training and inference overhead as well as the quality of reconstructed image in Fig. 8, the architecture with 1010 copies and 11 RB seems promising to implement in clinical scanners.

2.3 LSGAN for sharp MR images

Fig.8 compares the retrieved images by various recurrent GAN architectures with the input ZF image as well as the gold-standard one that is fully-sampled. Abdominal slices shown for two representative axial slices including liver and kidneys confirm again that RGANCS scheme with 1010 copies and 11 RB performs the best in terms of perceptual quality. Even though SNR and SSIM are not proper metrics to assess the perceptual quality, for the sake of completeness we report them in Table 2. This also corroborates even when using perceptual loss for training, recurrent scheme can significantly improve SNR/SSIM relative to a single deep network (11 copy, 1010 RBs) as commonly adopted for image restoration tasks in the literature. The RGANCS images are sharper than the CS-wavelet scheme, even though CS achieves a higher SNR/SSIM. Choosing a smaller weight λ\lambda, or, a larger η\eta, RGANCS can even improve the SNR/SSIM as it was seen in Table 1 of the paper. Further tunning of λ\lambda and η\eta for the best performance needs expert opinion of radiologists about the diagnostic quality of the resulting images and is the subject of our ongoing research.

3 Single image super-resolution

More evaluations are performed for super-resolving natural images. In essence, super-resolution can be seen as a linear inverse task, where one has only access to a low-resolution image y=ϕ∗x+v\mathbf{y}=\phi*\mathbf{x}+\mathbf{v} obtained after downsampling with a convolution kernel ϕ\phi. We adopt a 4×44\times 4 constant kernel with stride 44 that averages out the image pixel intensities over 4×44\times 4 non-overlapping regions. Image super-resolution is a challenging ill-posed problem, and has thus been the subject of intensive research over the last decade; see e.g., and the references therein. leverages deep convolutional GANs (DCGANs) accompanied with an affine projection layer to find a better solution as measured by SNR and SSIM. also deploys a deep ResNet (1616 RBs with 6464 feature maps) along with GAN perceptual cost to retrieve photo-realistic images. Our goal is not to create better looking images than the state-of-the-art, but to study proximal learning for this application that gives insights about possibly simpler network architectures for real-time tasks, and can interpret the proximal behavior in terms of revealing the details.

CelebA dataset. Adopting celebFaces Attributes Dataset (CelebA) , for training and test we use 1010k and 1,2801,280 images, respectively. Each ground-truth face image has 128×128128\times 128 pixels that is down-sampled to a 32×3232\times 32 low-resolution image.

Training. Our TensorFlow implementation uses a 2D conv with stride 44 for downsampling, and transpose conv with stride 44 for upsampling. Note, transpose convolution does not perform deconvolution, and a single conv transpose can be quite suboptimum. We approximate the deconvolution pseudo-inverse with a few (55) gradient-descent iterations with a small step size (0.10.1). The deconvolution then involves conv and transpose conv, and its gradient turns out to below up abruptly and generates NaNs during gradient backpropagation. We thus use gradient clipping for the G network that fixes the issue. The network is fed with pixelated 128×128128\times 128 images with three RGB channels obtained by an approximate deconvolution. The same network architecture as for MRI is adopted.

For the superresolution task when sharing the weights no interesting pattern is observed for ResNets of size 1−71-7 RBs, which indicates modeling the proximal needs larger networks. For the scenario with varying weights Fig. 9 plots PSNR and SSIM for various architectures. Using 55 copies with 22 RBs seems to perform as good as a deep ResNet with 1515 RBs adopted in . However, increasing the number of copies and RBs does not offer any clear advantages. This is in contrast with MRI reconstruction where a recurrent single RB could significantly outperform the deep architectures with up to 1010 RBs. It appears that recurrent architecture sounds more useful for global recovery tasks where the observation matrix entangles the image pixels. Perhaps by going to kk-space ResNet can better learn the proximals to invert the map. This is deferred to future research.

3.2 Interpretation of generator outputs

Fig. 10 depicts the output of different generator copies when training a network architecture with 44 independent copies, each composed of 55 RBs. RGB images are shown in gray scale. It is seen that different copies focus on features at different levels of abstraction associated with different frequency components. The first block tries to retrieve major (low-frequency) structural features at the expense of introducing a large amount of high-frequency noise, which is then washed away by the next copy. The third copy then adds up high-frequency components to improve the sharpness, which introduces some noise that is again alleviated by the fourth copy to retrieve the output. The overall process tend to alternate between sharpening and smoothing.

Conclusions and closing remarks

This paper caters a novel proximal learning framework for automated recovery of images from compressed linear measurements. Unrolling the proximal gradient iterations, a recurrent/cascade architecture is devised that alternates between proximal projection and data fidelity. ResNets are adopted to model the proximals, and a mixture of pixel-wise and perceptual costs used for training. Experiments are examined to assess various network wirings in reconstructing MR images of pediatric patients, and superresolving face images. Our observations indicate that a recurrent small ResNet can effectively learn the proximal, and significantly improve the quality and complexity of recent deep architectures (single copy) and the conventional CS-MRI. Our preliminary results for single-image suprerresolution however indicate that the recurrent architecture are not that effective compared with the exiting deep schemes.

There are still unanswered questions that are the focus of our current research. They pertain to running more experiments with perceptual costs with a subjective quality-assessment strategy; more extensive experiments for superresoltuion with larger ResNet sizes and possibly training in the kk-space with more global measurements; and a fair mechanism to compare the inference/training time.

References