Plug-and-play denoisers#

Open in Colab

This lesson regularizes an undersampled, noisy Cartesian SENSE reconstruction with a pretrained image denoiser in place of a specified penalty, and compares the result with total-variation regularization of the same data. The aim is to show how a denoiser enters a proximal iteration, what it improves on a hand-crafted penalty, and how its noise level plays the role of the regularization weight.

A proximal iteration such as ADMM or FISTA uses the regularization term only through its proximal operator,

\[\operatorname{prox}_{\gamma g}(v) = \arg\min_x \; \tfrac12 \|x - v\|_2^2 + \gamma\, g(x),\]

which is the maximum a posteriori estimate of an image \(x\) observed as \(v\) in white Gaussian noise, under the prior \(\exp(-g)\): the proximal operator is a denoiser. Plug-and-play regularization [1] [2] replaces it by any image denoiser \(D_\sigma\), without writing down \(g\). Each iteration alternates a step towards consistency with the measured k-space and a denoising step; the noise level \(\sigma\) of the denoiser takes the role of the regularization weight.

The denoiser is DRUNet [3], a convolutional network with the weights distributed by deepinv, trained for Gaussian denoising of natural grayscale photographs, not of MR images. bartorch.priors.ImplicitPrior converts between the complex image of the reconstruction and the real planes the network takes; the iterations are bartorch.optim.admm() and bartorch.optim.fista(), unchanged.

The phantom is the BrainWeb slice of Regularized reconstruction; the cell that builds it is hidden on this page and present in the downloadable script.

Learning objectives

  • Wrap a pretrained denoiser as bartorch.priors.ImplicitPrior and pass it to bartorch.optim.admm() and bartorch.optim.fista() in place of a bartorch.priors term.

  • Compare the result with total-variation regularization on the same data, in the images, the error maps and an enlarged region.

  • Vary the denoiser’s noise level and recognize under- and over-regularization.

It follows Parameter maps straight from k-space. The next lesson, MoDL, on BART’s ADMM, trains the denoiser through the iteration.

The pretrained weights, about 125 MB, are downloaded on the first call. The network runs once per iteration, which dominates the run time of this example on a CPU.

import csv
from pathlib import Path

import brainweb_dl
import numpy as np
import torch
from brainweb_dl import get_mri
from deepinv.models import DRUNet

import bartorch
import bartorch.tools as bt
from bartorch import linop, optim, priors

SIZE = 128
COILS = 8
ACCELERATION = 4
CALIBRATION = 16

Acquisition#

A quarter of the phase encodes (\(R = 4\)), drawn from a variable density around a fully sampled ACS region of 16 lines, with complex Gaussian noise of variance \(10^{-3}\) per sample of the unitary transform: the acquisition of Regularized reconstruction, where both noise amplification and incoherent aliasing limit an unregularized reconstruction. The sensitivities are the ones the data were simulated with, so that the comparison below concerns the regularization alone; Coil sensitivity calibration compares their estimation.

encodes = torch.arange(SIZE) - SIZE // 2
centre = (encodes.abs() < CALIBRATION // 2).to(torch.float32)
drawn = torch.multinomial(
    (1.0 + 2.0 * encodes.abs() / SIZE) ** -3.0 * (1.0 - centre),
    SIZE // ACCELERATION - CALIBRATION,
    replacement=False,
    generator=torch.Generator().manual_seed(11),
)
lines = centre.clone()
lines[drawn] = 1.0
pattern = lines.reshape(SIZE, 1).to(torch.complex64)

A = linop.CartesianSense(sensitivities, (SIZE, SIZE), pattern=pattern)
data = bt.noise(A(image), n=1e-3, s=42) * pattern

A specified regularizer#

Total variation under ADMM is the reference point, at the best of the weights 0.002, 0.005, 0.01 and 0.02 judged by NRMSE and SSIM against the phantom.

total_variation = optim.admm(data, A, priors.TotalVariation((-1, -2), 0.01), maxiter=60, rho=0.1)

A denoiser as the proximal step#

spatial=2 states that the network operates on two-dimensional planes of shape (n, channels, y, x). ImplicitPrior scales each image to unit peak modulus, denoises its real and imaginary parts as two grayscale planes, and scales the result back, so sigma is in units of the image’s peak. The network is evaluated without gradients, since nothing is trained here. Each iteration costs one application of the network.

ITERATIONS = 30

denoiser = DRUNet(in_channels=1, out_channels=1, pretrained="download").eval()
prior = priors.ImplicitPrior(denoiser, sigma=0.05, spatial=2)

with torch.no_grad():
    admm = optim.admm(data, A, prior, maxiter=ITERATIONS, rho=0.2)
    fista = optim.fista(data, A, prior, maxiter=ITERATIONS)

reconstructions = {
    "zero-filled": A.H(data),
    "total variation": total_variation,
    "DRUNet, ADMM": admm,
    "DRUNet, FISTA": fista,
}

for name, estimate in reconstructions.items():
    error = bt.nrmse(image.abs(), estimate.abs(), scaled=True)
    similarity = bt.ssim(image.abs(), scaled(estimate, image))
    print(f"{name:>16}  NRMSE {error:.3f}  SSIM {similarity:.3f}")
Downloading: "https://huggingface.co/deepinv/drunet/resolve/main/drunet_deepinv_gray_finetune_26k.pth?download=true" to /home/runner/.cache/torch/hub/checkpoints/drunet_deepinv_gray_finetune_26k.pth

  0%|          | 0.00/125M [00:00<?, ?B/s]
  0%|          | 128k/125M [00:00<06:17, 345kB/s]
  1%|          | 768k/125M [00:00<01:04, 2.02MB/s]
  2%|▏         | 2.50M/125M [00:00<00:19, 6.40MB/s]
  7%|▋         | 8.88M/125M [00:00<00:05, 23.2MB/s]
 23%|██▎       | 28.1M/125M [00:00<00:01, 73.6MB/s]
 30%|██▉       | 36.9M/125M [00:00<00:01, 73.3MB/s]
 36%|███▌      | 45.0M/125M [00:01<00:01, 68.4MB/s]
 42%|████▏     | 52.4M/125M [00:01<00:01, 68.2MB/s]
 48%|████▊     | 59.5M/125M [00:01<00:01, 65.2MB/s]
 53%|█████▎    | 66.1M/125M [00:01<00:01, 47.3MB/s]
 58%|█████▊    | 71.8M/125M [00:01<00:01, 49.6MB/s]
 64%|██████▍   | 79.8M/125M [00:01<00:00, 57.4MB/s]
 69%|██████▉   | 86.1M/125M [00:01<00:00, 59.7MB/s]
 74%|███████▍  | 92.4M/125M [00:02<00:00, 55.4MB/s]
 79%|███████▉  | 98.1M/125M [00:02<00:00, 54.1MB/s]
 84%|████████▍ | 105M/125M [00:02<00:00, 57.4MB/s]
 89%|████████▉ | 111M/125M [00:02<00:00, 54.9MB/s]
 93%|█████████▎| 116M/125M [00:02<00:00, 54.6MB/s]
 99%|█████████▉| 124M/125M [00:02<00:00, 60.0MB/s]
100%|██████████| 125M/125M [00:02<00:00, 50.5MB/s]
     zero-filled  NRMSE 0.199  SSIM 0.482
 total variation  NRMSE 0.161  SSIM 0.754
    DRUNet, ADMM  NRMSE 0.103  SSIM 0.893
   DRUNet, FISTA  NRMSE 0.135  SSIM 0.898
  • reference, zero-filled, total variation, DRUNet, ADMM
  • zero-filled error, total variation error, DRUNet, ADMM error, DRUNet, FISTA error
  • reference, enlarged, total variation, DRUNet, ADMM, DRUNet, FISTA

The zero-filled image shows the incoherent aliasing of the random sampling and the noise. Total variation removes most of both, but at its best weight it leaves a blotchy texture across the brain and flattens the gradual intensity variations into patches. The plug-and-play reconstruction under ADMM removes the noise and the aliasing while keeping the tissue boundaries, and its error map is darker inside the brain; in the enlarged region the ventricles and the larger cortical folds are delineated, although the finest sulci are lost. Under FISTA, at the same \(\sigma\), the same denoiser produces a much smoother image, in which the cortical folds have disappeared, and its error lies along every tissue boundary.

The two iterations differ in the step before the denoiser: FISTA takes a gradient step of fixed length on the data term, ADMM solves a quadratic problem that holds the image to the data with weight \(\rho\). The denoiser is therefore applied to different images, and its effective strength differs between the two iterations at the same \(\sigma\). A fixed \(\sigma\) makes neither iteration the minimization of a known objective, so the number of iterations and \(\rho\) enter the result, and are parameters to be chosen like the weight of a specified term. The printed SSIM ranks the over-smoothed FISTA image above the ADMM image: a single figure of merit does not replace looking at the images.

The noise level#

\(\sigma\) is the strength of the prior. A denoiser asked for less noise than the iterate contains leaves residual noise and aliasing in place; one asked for more removes image detail with them.

levels = (0.02, 0.05, 0.12)
with torch.no_grad():
    sweep = {
        sigma: admm
        if sigma == 0.05
        else optim.admm(
            data,
            A,
            priors.ImplicitPrior(denoiser, sigma=sigma, spatial=2),
            maxiter=ITERATIONS,
            rho=0.2,
        )
        for sigma in levels
    }

for sigma, estimate in sweep.items():
    error = bt.nrmse(image.abs(), estimate.abs(), scaled=True)
    similarity = bt.ssim(image.abs(), scaled(estimate, image))
    print(f"sigma {sigma:.2f}  NRMSE {error:.3f}  SSIM {similarity:.3f}")
sigma 0.02  NRMSE 0.139  SSIM 0.734
sigma 0.05  NRMSE 0.103  SSIM 0.893
sigma 0.12  NRMSE 0.127  SSIM 0.886
$\sigma$ = 0.02 (too small), $\sigma$ = 0.05, $\sigma$ = 0.12 (too large)

At the smallest \(\sigma\) the noise and the aliasing remain as a mottled texture; at the largest the cortex is smoothed into uniform white matter and small structures disappear; the printed errors have their minimum in between.

The denoiser was not trained on MR images, nor for the residual aliasing an undersampled acquisition leaves, which is not white Gaussian noise. The next lesson, MoDL, on BART’s ADMM, trains a network inside the iteration, on the acquisition it is applied to.

References#

Total running time of the script: (0 minutes 27.081 seconds)

Gallery generated by Sphinx-Gallery