
.. DO NOT EDIT.
.. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY.
.. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE:
.. "auto_examples/06-learning/01-plug-and-play.py"
.. LINE NUMBERS ARE GIVEN BELOW.

.. only:: html

    .. note::
        :class: sphx-glr-download-link-note

        :ref:`Go to the end <sphx_glr_download_auto_examples_06-learning_01-plug-and-play.py>`
        to download the full example code.

.. rst-class:: sphx-glr-example-title

.. _sphx_glr_auto_examples_06-learning_01-plug-and-play.py:


=======================
Plug-and-play denoisers
=======================

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,

.. math::

   \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 :math:`x` observed as
:math:`v` in white Gaussian noise, under the prior :math:`\exp(-g)`: the
proximal operator is a denoiser. Plug-and-play regularization
[#venkatakrishnan]_ [#ahmad]_ replaces it by any image denoiser
:math:`D_\sigma`, without writing down :math:`g`. Each iteration alternates a
step towards consistency with the measured k-space and a denoising step; the
noise level :math:`\sigma` of the denoiser takes the role of the
regularization weight.

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

The phantom is the BrainWeb slice of
:doc:`../03-regularization/01-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 :class:`bartorch.priors.ImplicitPrior` and
  pass it to :func:`bartorch.optim.admm` and :func:`bartorch.optim.fista` in
  place of a :mod:`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 :doc:`../05-model-based/02-quantitative-models`. The next lesson,
:doc:`02-modl-with-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.

.. GENERATED FROM PYTHON SOURCE LINES 59-254

.. code-block:: Python


    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









.. GENERATED FROM PYTHON SOURCE LINES 255-267

Acquisition
-----------

A quarter of the phase encodes (:math:`R = 4`), drawn from a variable
density around a fully sampled ACS region of 16 lines, with complex Gaussian
noise of variance :math:`10^{-3}` per sample of the unitary transform: the
acquisition of :doc:`../03-regularization/01-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; :doc:`../02-parallel-imaging/01-coil-calibration` compares their
estimation.

.. GENERATED FROM PYTHON SOURCE LINES 268-284

.. code-block:: Python


    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








.. GENERATED FROM PYTHON SOURCE LINES 285-291

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.

.. GENERATED FROM PYTHON SOURCE LINES 292-295

.. code-block:: Python


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








.. GENERATED FROM PYTHON SOURCE LINES 296-305

A denoiser as the proximal step
-------------------------------

``spatial=2`` states that the network operates on two-dimensional planes of
shape ``(n, channels, y, x)``. :class:`~bartorch.priors.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.

.. GENERATED FROM PYTHON SOURCE LINES 306-328

.. code-block:: Python


    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}")





.. rst-class:: sphx-glr-script-out

 .. code-block:: none

    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




.. GENERATED FROM PYTHON SOURCE LINES 329-361




.. rst-class:: sphx-glr-horizontal


    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_001.png
         :alt: reference, zero-filled, total variation, DRUNet, ADMM
         :srcset: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_001.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_002.png
         :alt: zero-filled error, total variation error, DRUNet, ADMM error, DRUNet, FISTA error
         :srcset: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_002.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_003.png
         :alt: reference, enlarged, total variation, DRUNet, ADMM, DRUNet, FISTA
         :srcset: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_003.png
         :class: sphx-glr-multi-img





.. GENERATED FROM PYTHON SOURCE LINES 362-390

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 :math:`\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 :math:`\rho`. The
denoiser is therefore applied to different images, and its effective
strength differs between the two iterations at the same :math:`\sigma`. A
fixed :math:`\sigma` makes neither iteration the minimization of a known
objective, so the number of iterations and :math:`\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
---------------

:math:`\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.

.. GENERATED FROM PYTHON SOURCE LINES 391-412

.. code-block:: Python


    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}")





.. rst-class:: sphx-glr-script-out

 .. code-block:: none

    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




.. GENERATED FROM PYTHON SOURCE LINES 413-422




.. image-sg:: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_004.png
   :alt: $\sigma$ = 0.02 (too small), $\sigma$ = 0.05, $\sigma$ = 0.12 (too large)
   :srcset: /auto_examples/06-learning/images/sphx_glr_01-plug-and-play_004.png
   :class: sphx-glr-single-img





.. GENERATED FROM PYTHON SOURCE LINES 423-432

At the smallest :math:`\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, :doc:`02-modl-with-admm`, trains a network inside the iteration, on
the acquisition it is applied to.

.. GENERATED FROM PYTHON SOURCE LINES 435-451

References
----------

.. [#venkatakrishnan] Venkatakrishnan SV, Bouman CA, Wohlberg B. Plug-and-play
   priors for model based reconstruction. *IEEE Global Conference on Signal and
   Information Processing*, 945-948 (2013).
   https://doi.org/10.1109/GlobalSIP.2013.6737048

.. [#ahmad] Ahmad R, Bouman CA, Buzzard GT, Chan S, Liu S, Reehorst ET,
   Schniter P. Plug-and-play methods for magnetic resonance imaging: using
   denoisers for image recovery. *IEEE Signal Process Mag* 37(1):105-116
   (2020). https://doi.org/10.1109/MSP.2019.2949470

.. [#zhang] Zhang K, Li Y, Zuo W, Zhang L, Van Gool L, Timofte R. Plug-and-play
   image restoration with deep denoiser prior. *IEEE Trans Pattern Anal Mach
   Intell* 44(10):6360-6376 (2022). https://doi.org/10.1109/TPAMI.2021.3088914


.. rst-class:: sphx-glr-timing

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


.. _sphx_glr_download_auto_examples_06-learning_01-plug-and-play.py:

.. only:: html

  .. container:: sphx-glr-footer sphx-glr-footer-example

    .. container:: sphx-glr-download sphx-glr-download-jupyter

      :download:`Download Jupyter notebook: 01-plug-and-play.ipynb <01-plug-and-play.ipynb>`

    .. container:: sphx-glr-download sphx-glr-download-python

      :download:`Download Python source code: 01-plug-and-play.py <01-plug-and-play.py>`

    .. container:: sphx-glr-download sphx-glr-download-zip

      :download:`Download zipped: 01-plug-and-play.zip <01-plug-and-play.zip>`


.. only:: html

 .. rst-class:: sphx-glr-signature

    `Gallery generated by Sphinx-Gallery <https://sphinx-gallery.github.io>`_
