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

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

.. _sphx_glr_auto_examples_06-learning_05-self-supervised-training.py:


============================
Training without a reference
============================

**Aim.** Train the unrolled network of :doc:`04-staged-training` from
undersampled k-space alone, with no fully sampled reference, and measure how
much of the supervised network's image quality it retains.

Dynamic and high-dimensional acquisitions -- a cine, a functional run, a
fingerprinting series -- are rarely acquired fully sampled: undersampling is
what makes them feasible within a breath-hold or a scan time. There is then no
reference image to train against. Self-supervised learning via data
undersampling (SSDU) [#ssdu]_ trains against the measured k-space itself. The
acquired phase encodes :math:`\Omega` are split into two disjoint sets,
:math:`\Theta` and :math:`\Lambda`; the network reconstructs from
:math:`\Theta`, and the loss compares the k-space of its reconstruction with
the measured data on the held-out set :math:`\Lambda`,

.. math::

   \mathcal{L} = \frac{\|y_\Lambda - A_\Lambda f_\theta(y_\Theta)\|_2}{\|y_\Lambda\|_2}
   + \frac{\|y_\Lambda - A_\Lambda f_\theta(y_\Theta)\|_1}{\|y_\Lambda\|_1},

where :math:`A_\Lambda` is the SENSE encoding restricted to :math:`\Lambda`.
A new split is drawn at every step, so over the training every acquired line
is both reconstructed from and held out [#multimask]_. At inference the
network reconstructs from all of :math:`\Omega`.

**Learning objectives**

- Partition the acquired phase encodes with :func:`bartorch.learning.split`.
- Train an unrolled network self-supervised with
  :class:`bartorch.learning.Reconstruction`, by giving items the
  sampling pattern instead of a reference.
- Compare with the same network trained against references, and with
  CG-SENSE.

It follows :doc:`04-staged-training`. The next lesson,
:doc:`06-annealed-plug-and-play`, uses a denoiser trained once for any
acquisition.

.. GENERATED FROM PYTHON SOURCE LINES 45-154

.. 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, SSIMMetric
    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 = 4
    ACCELERATION = 4
    EPOCHS = 16

    _ = torch.manual_seed(0)








.. GENERATED FROM PYTHON SOURCE LINES 155-161

Data
----

The slices, coils and fourfold undersampling of :doc:`04-staged-training`:
subject 0 to train on and subject 4 to validate on. The references are kept
only to score the results; the self-supervised network never sees them.

.. GENERATED FROM PYTHON SOURCE LINES 162-243

.. 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)

    density = torch.exp(-0.5 * ((torch.arange(SIZE) - SIZE / 2) / (SIZE / 6)) ** 2)
    lines = torch.rand(SIZE, generator=torch.Generator().manual_seed(1)) < density / density.sum() * (
        SIZE / ACCELERATION
    )
    lines[SIZE // 2 - 4 : SIZE // 2 + 4] = True
    pattern = lines.to(torch.complex64)[:, None].expand(SIZE, SIZE).contiguous()
    A = linop.CartesianSense(sensitivities, (SIZE, SIZE), pattern=pattern)
    NOISE = 0.02

    generator = torch.Generator().manual_seed(3)
    kspace = {
        "train": [
            A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
            for x in train_images
        ],
        "valid": [
            A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
            for x in valid_images
        ],
    }








.. GENERATED FROM PYTHON SOURCE LINES 244-254

The split
---------

The readout is fully sampled, so the unit of the split is the phase-encode
line: the pattern given to :func:`~bartorch.learning.split` has one entry per
line and broadcasts over the coils and the readout. A quarter of the acquired
lines are held out, drawn with a Gaussian density across :math:`k_y`, and
the eight central lines always stay in :math:`\Theta`: a reconstruction
without the centre of k-space would lose the image contrast, and the loss
would be dominated by it.

.. GENERATED FROM PYTHON SOURCE LINES 255-279

.. code-block:: Python


    acquired = lines.float()[:, None]
    keep, held = learning.split(acquired, 0.25, keep=(8, 1), generator=torch.Generator().manual_seed(0))
    print(
        f"{int(acquired.sum())} acquired lines: {int(keep.sum())} to reconstruct from, "
        f"{int(held.sum())} held out"
    )





.. image-sg:: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_001.png
   :alt: 05 self supervised training
   :srcset: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_001.png
   :class: sphx-glr-single-img


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

 .. code-block:: none

    24 acquired lines: 20 to reconstruct from, 4 held out




.. GENERATED FROM PYTHON SOURCE LINES 280-290

Two networks, one trained each way
----------------------------------

Both are the iteration-conditioned unrolled network of the previous lesson,
trained end to end for the same number of epochs from the same
initialization; only the items differ. A self-supervised item carries the
acquired ``pattern`` in place of a ``target``, and
:class:`~bartorch.learning.Reconstruction` then draws a split at
every step. Its validation loss is the held-out loss on a split fixed for
the whole run, and needs no reference either.

.. GENERATED FROM PYTHON SOURCE LINES 291-331

.. code-block:: Python



    def unrolled():
        torch.manual_seed(0)
        network = learning.UNet(2, spatial=2, widths=(16, 32, 64), steps=True)
        prior = priors.ImplicitPrior(learning.ComplexNet(network, spatial=2), step=True)
        block = optim.ISTBlock(prior, step=1.0)
        block.step.requires_grad_()
        return learning.Unrolled(block, iterations=ITERATIONS, checkpoint=True)


    def items(part, supervised):
        images = train_images if "train" == part else valid_images
        made = []
        for x, y in zip(images, kspace[part]):
            item = {"y": y, "A": A}
            item.update({"target": x} if supervised else {"pattern": acquired})
            made.append(item)
        return made


    models = {}
    for name, supervised in (("supervised", True), ("self-supervised", False)):
        models[name] = unrolled()
        trainer = lightning.Trainer(
            max_epochs=EPOCHS,
            accelerator="cpu",
            logger=False,
            enable_checkpointing=False,
            enable_model_summary=False,
            enable_progress_bar=False,
        )
        trainer.fit(
            learning.Reconstruction(
                models[name], "end-to-end", lr=1e-3, fraction=0.25, split_options={"keep": (8, 1)}
            ),
            DataLoader(items("train", supervised), batch_size=4, shuffle=True, collate_fn=list),
            DataLoader(items("valid", supervised), batch_size=4, collate_fn=list),
        )





.. 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.




.. GENERATED FROM PYTHON SOURCE LINES 332-337

Results
-------

Both networks reconstruct from all acquired lines of the validation subject
and are scored against its reference.

.. GENERATED FROM PYTHON SOURCE LINES 338-355

.. code-block:: Python


    psnr = PSNRMetric(max_val=1.0)
    ssim = SSIMMetric(spatial_dims=2, data_range=1.0)
    truth = torch.stack(valid_images).abs()[:, None]

    with torch.no_grad():
        results = {name: torch.stack([m(y, A) for y in kspace["valid"]]) for name, m in models.items()}
    results["CG SENSE, 20 iterations"] = torch.stack(
        [optim.cg(y, A, maxiter=20) for y in kspace["valid"]]
    )
    for name, made in results.items():
        a = made.abs()[:, None]
        print(
            f"{name:>24}   PSNR {float(psnr(a, truth).mean()):5.2f} dB   "
            f"SSIM {float(ssim(a, truth).mean()):.3f}"
        )





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

 .. code-block:: none

                  supervised   PSNR 30.20 dB   SSIM 0.945
             self-supervised   PSNR 28.59 dB   SSIM 0.726
     CG SENSE, 20 iterations   PSNR 24.27 dB   SSIM 0.558




.. GENERATED FROM PYTHON SOURCE LINES 356-368

The self-supervised network is trained on less information: each step
reconstructs from three quarters of the acquired lines and is told nothing
about the lines never acquired, which is where the supervised network learns
most. The gap between the two is what a fully sampled reference would have
bought; the self-supervised network needs nothing beyond the data the
protocol already acquires.

In the images below, both networks suppress the noise that the CG-SENSE
unfolding amplifies across the whole field of view. The self-supervised
network keeps more residual aliasing along the phase-encode direction
(vertical), which its error map shows as horizontal striping: the lines never
acquired are the ones it cannot score against.

.. GENERATED FROM PYTHON SOURCE LINES 369-383




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


    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_002.png
         :alt: reference, CG-SENSE, supervised, self-supervised
         :srcset: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_002.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_003.png
         :alt: reference, enlarged, CG-SENSE, enlarged, supervised, enlarged, self-supervised, enlarged
         :srcset: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_003.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_004.png
         :alt: CG-SENSE NRMSE 0.124, supervised NRMSE 0.063, self-supervised NRMSE 0.076
         :srcset: /auto_examples/06-learning/images/sphx_glr_05-self-supervised-training_004.png
         :class: sphx-glr-multi-img





.. GENERATED FROM PYTHON SOURCE LINES 384-396

References
----------

.. [#ssdu] Yaman B, Hosseini SAH, Moeller S, Ellermann J, Ugurbil K, Akcakaya M.
   Self-supervised learning of physics-guided reconstruction neural networks
   without fully sampled reference data. *Magn Reson Med* 84(6):3172-3191
   (2020). https://doi.org/10.1002/mrm.28378

.. [#multimask] Yaman B, Hosseini SAH, Moeller S, Ellermann J, Ugurbil K,
   Akcakaya M. Multi-mask self-supervised learning for physics-guided neural
   networks in highly accelerated magnetic resonance imaging. *NMR Biomed*
   35(12):e4798 (2022). https://doi.org/10.1002/nbm.4798


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

   **Total running time of the script:** (2 minutes 50.589 seconds)


.. _sphx_glr_download_auto_examples_06-learning_05-self-supervised-training.py:

.. only:: html

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

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

      :download:`Download Jupyter notebook: 05-self-supervised-training.ipynb <05-self-supervised-training.ipynb>`

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

      :download:`Download Python source code: 05-self-supervised-training.py <05-self-supervised-training.py>`

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

      :download:`Download zipped: 05-self-supervised-training.zip <05-self-supervised-training.zip>`


.. only:: html

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

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