
.. DO NOT EDIT.
.. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY.
.. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE:
.. "auto_examples/06-learning/06-annealed-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_06-annealed-plug-and-play.py>`
        to download the full example code.

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

.. _sphx_glr_auto_examples_06-learning_06-annealed-plug-and-play.py:


===========================
Annealed plug-and-play
===========================

**Aim.** Train one denoiser on images alone, without any encoding, and use it
as the regularizer of an ADMM reconstruction at any undersampling, with the
denoising strength decreasing over the iterations; show that it holds up at
an acceleration where CG-SENSE breaks down.

A plug-and-play reconstruction [#pnp]_ replaces the proximal step of an
iterative algorithm by a denoiser. The denoiser is trained on images, not on
k-space, so one network serves every protocol whose images resemble its
training images: a change of acceleration, sampling pattern or coil array
needs no retraining. The proximal step of a regularizer :math:`\lambda\phi`
with ADMM penalty :math:`\rho` is a Gaussian denoiser of noise variance
:math:`\sigma^2 = \lambda/\rho`. The first iterates carry the strong
incoherent aliasing of the undersampling and the last are nearly consistent
with the data, so a noise level that decreases from one to the other
[#dpir]_ removes the aliasing first and preserves fine anatomy at the end. The
penalty follows as :math:`\rho_k = \lambda/\sigma_k^2`, which keeps the
balance between data consistency and denoising that :math:`\lambda` sets.

**Learning objectives**

- Train a denoiser conditioned on the noise level, ``noise=True`` in
  :class:`bartorch.learning.UNet`, with
  :class:`bartorch.learning.Reconstruction`.
- Give :class:`bartorch.priors.ImplicitPrior` a schedule of noise levels and
  :class:`bartorch.optim.ADMMBlock` the matching schedule of penalties.
- Compare an annealed schedule with a fixed noise level, iteration by
  iteration, and apply the same denoiser at a higher acceleration.

It follows :doc:`05-self-supervised-training`. The next lesson,
:doc:`07-uncertainty`, attaches error bars to a learned reconstruction.

.. GENERATED FROM PYTHON SOURCE LINES 38-145

.. code-block:: Python


    import csv
    import logging
    from pathlib import Path

    import brainweb_dl
    import lightning
    import numpy as np
    import torch
    from brainweb_dl import get_mri
    from monai.metrics import PSNRMetric
    from torch.utils.data import DataLoader

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

    SIZE = 96
    COILS = 8
    ITERATIONS = 12

    _ = torch.manual_seed(0)








.. GENERATED FROM PYTHON SOURCE LINES 146-151

Data
----

The slices, coils and fourfold undersampling of :doc:`04-staged-training`:
subject 0 to train the denoiser on, subject 4 to reconstruct.

.. GENERATED FROM PYTHON SOURCE LINES 152-234

.. code-block:: Python


    train_images = brain_slices(subject=0, count=24)
    valid_images = brain_slices(subject=4, count=8)

    sensitivities = bt.coils(t=bt.grid(D=(SIZE, SIZE, 1)), n=COILS)[:, 0]
    sensitivities = sensitivities / bartorch.rss(sensitivities, axes=(0,), keepdim=True)
    NOISE = 0.02


    def acquisition(acceleration, seed):
        """A Cartesian SENSE encoding sampling one line in ``acceleration`` on average."""
        density = torch.exp(-0.5 * ((torch.arange(SIZE) - SIZE / 2) / (SIZE / 6)) ** 2)
        chance = density / density.sum() * (SIZE / acceleration)
        lines = torch.rand(SIZE, generator=torch.Generator().manual_seed(seed)) < chance
        lines[SIZE // 2 - 4 : SIZE // 2 + 4] = True
        pattern = lines.to(torch.complex64)[:, None].expand(SIZE, SIZE).contiguous()
        return linop.CartesianSense(sensitivities, (SIZE, SIZE), pattern=pattern)


    def measure(A, generator):
        return [
            A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
            for x in valid_images
        ]


    A = acquisition(4, seed=1)
    kspace = measure(A, torch.Generator().manual_seed(3))








.. GENERATED FROM PYTHON SOURCE LINES 235-245

A denoiser for every noise level
--------------------------------

With ``noise=True`` the U-Net takes the noise level as an input and
modulates its features with it, so one network covers a range of SNRs. It is
trained on pairs of a slice and the same slice with complex white Gaussian
noise, at a standard deviation drawn log-uniformly between 0.5 and 20 per
cent of the image's peak for each pair
(:class:`~bartorch.learning.ComplexNet` scales each image to unit peak). No
coil sensitivities, sampling pattern or k-space enter the training.

.. GENERATED FROM PYTHON SOURCE LINES 246-293

.. code-block:: Python


    LOW, HIGH = 0.005, 0.2

    network = learning.UNet(2, spatial=2, widths=(16, 32, 64), noise=True)
    denoiser = learning.ComplexNet(network, spatial=2)


    def pairs(images):
        made = []
        for x in images:
            sigma = LOW * (HIGH / LOW) ** torch.rand(())
            noisy = x + sigma * torch.randn_like(x)
            made.append({"input": noisy, "target": x, "sigma": sigma.reshape(1)})
        return made


    trainer = lightning.Trainer(
        max_epochs=30,
        accelerator="cpu",
        logger=False,
        enable_checkpointing=False,
        enable_model_summary=False,
        enable_progress_bar=False,
    )
    trainer.fit(
        learning.Reconstruction(denoiser, "denoiser", lr=2e-3),
        DataLoader(train_images, batch_size=4, shuffle=True, collate_fn=pairs),
        DataLoader(valid_images, batch_size=4, collate_fn=pairs),
    )

    psnr = PSNRMetric(max_val=1.0)
    truth = torch.stack(valid_images).abs()[:, None]


    def score(images):
        return float(psnr(images.abs()[:, None], truth).mean())


    with torch.no_grad():
        for sigma in (0.01, 0.05, 0.1):
            noisy = torch.stack(valid_images)
            noisy = noisy + sigma * torch.randn_like(noisy)
            denoised = torch.stack([denoiser(x[None], torch.tensor([sigma]))[0] for x in noisy])
            print(
                f"sigma {sigma:4.2f}: noisy {score(noisy):5.2f} dB, denoised {score(denoised):5.2f} dB"
            )





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

 .. code-block:: none

    /opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/utilities/_pytree.py:21: `isinstance(treespec, LeafSpec)` is deprecated, use `isinstance(treespec, TreeSpec) and treespec.is_leaf()` instead.
    /opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/trainer/connectors/data_connector.py:434: The 'val_dataloader' does not have many workers which may be a bottleneck. Consider increasing the value of the `num_workers` argument` to `num_workers=3` in the `DataLoader` to improve performance.
    /opt/hostedtoolcache/Python/3.12.14/x64/lib/python3.12/site-packages/lightning/pytorch/trainer/connectors/data_connector.py:434: The 'train_dataloader' does not have many workers which may be a bottleneck. Consider increasing the value of the `num_workers` argument` to `num_workers=3` in the `DataLoader` to improve performance.
    sigma 0.01: noisy 41.46 dB, denoised 41.45 dB
    sigma 0.05: noisy 27.17 dB, denoised 27.17 dB
    sigma 0.10: noisy 21.06 dB, denoised 21.06 dB




.. GENERATED FROM PYTHON SOURCE LINES 294-307

Schedules of noise level and penalty
------------------------------------

:class:`~bartorch.priors.ImplicitPrior` takes a sequence of noise levels,
one per iteration, and :class:`~bartorch.optim.ADMMBlock` a sequence of
penalties; the last value of each is repeated past its end. The annealed
schedule decreases the noise level geometrically from 0.1 to 0.01 of the
peak over twelve iterations, and the penalty follows as
:math:`\rho_k = \lambda/\sigma_k^2`. When the penalty changes, the block
rescales ADMM's scaled dual variable by the ratio of the old penalty to the
new, so that the unscaled one carries over. The two fixed schedules hold the
noise level at either end of the annealed one, with the penalty given by the
same :math:`\lambda`.

.. GENERATED FROM PYTHON SOURCE LINES 308-347

.. code-block:: Python


    LAMBDA = 1e-4

    schedules = {
        "annealed, 0.1 to 0.01": torch.logspace(-1.0, -2.0, ITERATIONS),
        "fixed, 0.03": torch.full((ITERATIONS,), 0.03),
        "fixed, 0.01": torch.full((ITERATIONS,), 0.01),
    }


    def plug_and_play(sigma):
        prior = priors.ImplicitPrior(denoiser, sigma=sigma.tolist())
        block = optim.ADMMBlock(prior, rho=(LAMBDA / sigma**2).tolist(), cg_maxiter=5)
        return learning.Unrolled(block, iterations=ITERATIONS)


    curves, results = {}, {}
    with torch.no_grad():
        for name, sigma in schedules.items():
            runs = [list(plug_and_play(sigma).steps(y, A)) for y in kspace]
            iterates = [torch.stack([run[k] for run in runs]) for k in range(ITERATIONS)]
            curves[name] = [score(images) for images in iterates]
            results[name] = iterates[-1]
        results["CG SENSE, 20 iterations"] = torch.stack([optim.cg(y, A, maxiter=20) for y in kspace])

    for name, images in results.items():
        print(f"{name:>24}   PSNR {score(images):5.2f} dB")





.. image-sg:: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_001.png
   :alt: 06 annealed plug and play
   :srcset: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_001.png
   :class: sphx-glr-single-img


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

 .. code-block:: none

       annealed, 0.1 to 0.01   PSNR 27.02 dB
                 fixed, 0.03   PSNR 27.14 dB
                 fixed, 0.01   PSNR 27.50 dB
     CG SENSE, 20 iterations   PSNR 24.21 dB




.. GENERATED FROM PYTHON SOURCE LINES 348-355

A large fixed noise level converges within a few iterations, to an image
limited by the smoothing the denoiser applies at that level. A small one
preserves detail, but its large penalty makes each ADMM step move little
from the previous one, and twelve iterations do not reach its fixed point.
The annealed schedule takes the large steps first and the small ones last,
and ends slightly ahead of the better fixed level without that level having
to be tuned for the acquisition.

.. GENERATED FROM PYTHON SOURCE LINES 356-369




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


    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_002.png
         :alt: reference, CG-SENSE, annealed plug-and-play
         :srcset: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_002.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_003.png
         :alt: CG-SENSE NRMSE 0.124, annealed plug-and-play NRMSE 0.090
         :srcset: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_003.png
         :class: sphx-glr-multi-img





.. GENERATED FROM PYTHON SOURCE LINES 370-378

Another acquisition, the same denoiser
--------------------------------------

The denoiser was trained without an encoding, so it applies unchanged to
:math:`R = 6`, a sampling pattern it has never been used with. At this
acceleration eight coils no longer unfold the aliasing well: CG-SENSE is
dominated by g-factor noise and residual aliasing, while the plug-and-play
reconstruction keeps the anatomy.

.. GENERATED FROM PYTHON SOURCE LINES 379-398

.. code-block:: Python


    A6 = acquisition(6, seed=2)
    kspace6 = measure(A6, torch.Generator().manual_seed(4))
    with torch.no_grad():
        annealed = torch.stack(
            [plug_and_play(schedules["annealed, 0.1 to 0.01"])(y, A6) for y in kspace6]
        )
        cg = torch.stack([optim.cg(y, A6, maxiter=20) for y in kspace6])
    print(f"sixfold: annealed plug-and-play {score(annealed):5.2f} dB, CG SENSE {score(cg):5.2f} dB")





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


    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_004.png
         :alt: reference, CG-SENSE, R = 6, annealed, R = 6
         :srcset: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_004.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_005.png
         :alt: CG-SENSE, R = 6 NRMSE 0.168, annealed, R = 6 NRMSE 0.139
         :srcset: /auto_examples/06-learning/images/sphx_glr_06-annealed-plug-and-play_005.png
         :class: sphx-glr-multi-img


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

 .. code-block:: none

    sixfold: annealed plug-and-play 23.72 dB, CG SENSE 21.98 dB




.. GENERATED FROM PYTHON SOURCE LINES 399-410

References
----------

.. [#pnp] 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

.. [#dpir] 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.526 seconds)


.. _sphx_glr_download_auto_examples_06-learning_06-annealed-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: 06-annealed-plug-and-play.ipynb <06-annealed-plug-and-play.ipynb>`

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

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

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

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


.. only:: html

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

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