
.. DO NOT EDIT.
.. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY.
.. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE:
.. "auto_examples/06-learning/03-networks-for-complex-volumes.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_03-networks-for-complex-volumes.py>`
        to download the full example code.

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

.. _sphx_glr_auto_examples_06-learning_03-networks-for-complex-volumes.py:


===============================
Networks for complex volumes
===============================

**Aim.** Train a 3D convolutional denoiser on patches of a complex,
multi-contrast brain volume, apply it to a whole volume of another subject
patch by patch, as it would run on a scanner GPU too small for the volume,
and check that the patch boundaries leave no visible seams.

The networks of the previous lessons denoise a single complex 2D slice. The
data a learned reconstruction is most needed for are larger: a 3D volume of
several contrasts, of subspace coefficients in MR fingerprinting, or of the
frames of a cine. Three things change. The contrasts are denoised jointly, as
channels of one image, so that the network can use the anatomy they share;
their signal levels differ, so they are balanced before the network sees
them. The volume does not fit the network's activations in GPU memory, so the
network is trained on patches and applied patch by patch. And a network
applied on a fixed grid of patches leaves seams at the patch boundaries,
which a random offset of the grid averages out.

**Learning objectives**

- Build a 3D :class:`bartorch.learning.UNet` for complex multi-contrast images
  with :class:`bartorch.learning.ComplexNet`, and balance the contrasts by
  whitening.
- Train it on patches drawn by ``torchio``, with augmentations that preserve
  the complex MR signal.
- Apply it to a whole volume with :class:`bartorch.learning.Patchwise`, and
  average the seams out with :func:`bartorch.learning.moments`.
- Compare the size of spatial and spatiotemporal networks.

It follows :doc:`02-modl-with-admm`. The next lesson,
:doc:`04-staged-training`, trains an unrolled network in stages.

.. GENERATED FROM PYTHON SOURCE LINES 38-143

.. code-block:: Python


    import csv
    import logging
    from pathlib import Path

    import brainweb_dl
    import lightning
    import numpy as np
    import torch
    import torchio as tio
    from brainweb_dl import get_mri
    from torch.utils.data import DataLoader

    from bartorch import learning

    SIZE = 64
    PATCH = 32
    NOISE = 0.06

    _ = torch.manual_seed(0)








.. GENERATED FROM PYTHON SOURCE LINES 144-155

A multi-contrast complex volume
-------------------------------

Three spin-echo contrasts of a BrainWeb subject -- :math:`T_1`-weighted
(TR 600 ms, TE 12 ms), :math:`T_2`-weighted (TR 4000 ms, TE 100 ms) and
proton-density-weighted (TR 4000 ms, TE 12 ms) -- on a :math:`64^3` grid,
each with its own smooth background phase, as a ``(3, z, y, x)`` complex
tensor. Each voxel's signal is the sum of the spin-echo signals of the
tissues it contains. Subject 0 is the training volume and subject 4 the test
volume. Complex Gaussian noise of 6 per cent of each contrast's peak gives
the test volume an SNR typical of a fast high-resolution scan.

.. GENERATED FROM PYTHON SOURCE LINES 156-209

.. code-block:: Python


    train_volume = brain_volume(subject=0)
    test_volume = brain_volume(subject=4)
    noisy = test_volume + NOISE * torch.randn_like(test_volume)
    print(f"test volume {tuple(test_volume.shape)}, {test_volume.dtype}")





.. image-sg:: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_001.png
   :alt: T1w, T2w, PDw, T1w, noisy, T2w, noisy, PDw, noisy
   :srcset: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_001.png
   :class: sphx-glr-single-img


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

 .. code-block:: none

    test volume (3, 64, 64, 64), torch.complex64




.. GENERATED FROM PYTHON SOURCE LINES 210-222

The network
-----------

:class:`~bartorch.learning.UNet` is a residual U-Net whose last convolution
starts at zero, so that the untrained network returns its input.
:class:`~bartorch.learning.ComplexNet` lays the three complex contrasts out
as its six real channels -- the real parts, then the imaginary parts -- and
lays its output back out as complex contrasts. With
``normalize="whiten"`` it subtracts each channel's mean and multiplies by the
inverse square root of the channels' covariance before the call, and undoes
both after it, so that the network sees uncorrelated channels of unit
variance whatever the relative energy of the contrasts.

.. GENERATED FROM PYTHON SOURCE LINES 223-235

.. code-block:: Python


    net = learning.UNet(6, spatial=3, widths=(8, 16, 32))
    denoiser = learning.ComplexNet(net, spatial=3, channels=1, normalize="whiten")
    weights = sum(p.numel() for p in net.parameters())
    print(f"{weights} weights, {4 * weights / 1e6:.2f} MB in single precision")

    with torch.no_grad():
        print(
            "untrained network returns its input:",
            torch.allclose(denoiser(noisy[None])[0], noisy, atol=1e-5),
        )





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

 .. code-block:: none

    136942 weights, 0.55 MB in single precision
    untrained network returns its input: True




.. GENERATED FROM PYTHON SOURCE LINES 236-253

Patches and augmentation
------------------------

``torchio`` holds the training volume as a :class:`torchio.ScalarImage` of
real channels (:func:`~bartorch.learning.as_real`) and draws :math:`32^3`
patches from it through a :class:`torchio.Queue`. The augmentations are
those that map one MR image to another the acquisition could have produced:
a flip, and :class:`~bartorch.learning.RandomGain`, a receiver gain
and global phase shared by the contrasts. An intensity transform applied to
the real and imaginary channels separately, such as a gamma correction,
would produce a signal no acquisition can. The global phase is varied over a
limited range: a network trained over every phase must learn to commute with
a rotation of its real and imaginary channels, which takes more training
than this lesson runs.

Each patch becomes a training pair when a new draw of noise is added to it,
so the network sees a different noise realization at every epoch.

.. GENERATED FROM PYTHON SOURCE LINES 254-291

.. code-block:: Python


    subject = tio.Subject(image=tio.ScalarImage(tensor=learning.as_real(train_volume).flatten(0, 1)))
    augment = tio.Compose(
        [tio.RandomFlip(axes=(0, 1, 2)), learning.RandomGain(phase=0.3, log_scale=0.2)]
    )
    queue = tio.Queue(
        tio.SubjectsDataset([subject], transform=augment),
        max_length=64,
        samples_per_volume=32,
        sampler=tio.UniformSampler(PATCH),
        num_workers=0,
    )


    def pairs(patches):
        made = []
        for patch in patches:
            clean = learning.as_complex(patch["image"][tio.DATA].unflatten(0, (2, 3)))
            made.append({"input": clean + NOISE * torch.randn_like(clean), "target": clean})
        return made


    validation = [{"input": noisy, "target": test_volume}]
    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(queue, batch_size=4, collate_fn=pairs),
        DataLoader(validation, batch_size=1, 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 292-312

Applying the network patch by patch
-----------------------------------

:class:`~bartorch.learning.Patchwise` keeps the volume in host memory and
sends it to the network's device a few patches at a time, runs the network
there in mixed precision, and assembles the result on the host. On a GPU
only the network and ``batch`` patches are resident, so a volume larger than
the GPU memory -- a whole-brain fingerprinting series on a scanner's
16 GB GPU -- is denoised by a network trained on patches of it. At
inference the copies of one group of patches overlap the computation on the
previous one. Here, on the host, the whole volume is also small enough to be
denoised in one call, which is the reference the patchwise result is
compared with.

A U-Net is not translation invariant at a patch boundary: its receptive field
extends past the patch, where it sees zeros rather than the neighbouring
voxels. On a fixed grid of patches the errors fall along the same planes
every time. With ``shift=True``, the grid is offset at random at every call,
and averaging a few calls with :func:`~bartorch.learning.moments` spreads
the boundary errors across the volume.

.. GENERATED FROM PYTHON SOURCE LINES 313-359

.. code-block:: Python


    whole_net = learning.ComplexNet(net, spatial=3, channels=1, normalize="whiten")
    fixed = learning.ComplexNet(
        learning.Patchwise(net, (PATCH,) * 3, shift=False), spatial=3, channels=1, normalize="whiten"
    )
    shifted = learning.ComplexNet(
        learning.Patchwise(net, (PATCH,) * 3, shift=True), spatial=3, channels=1, normalize="whiten"
    )

    with torch.no_grad():
        whole = whole_net(noisy[None])[0]
        grid = fixed(noisy[None])[0]
        averaged, spread = learning.moments(lambda: shifted(noisy[None])[0], samples=8)


    def error(made):
        return float((made - test_volume).norm() / test_volume.norm())


    print(f"noisy                  relative error {error(noisy):.4f}")
    for name, made in (("whole volume", whole), ("fixed grid", grid), ("8 random grids", averaged)):
        print(
            f"{name:<22} relative error {error(made):.4f}, "
            f"departure from the whole volume {float((made - whole).norm() / whole.norm()):.4f}"
        )





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


    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_002.png
         :alt: reference, noisy, denoised, whole
         :srcset: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_002.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_003.png
         :alt: noisy NRMSE 0.131, denoised, whole NRMSE 0.098
         :srcset: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_003.png
         :class: sphx-glr-multi-img

    *

      .. image-sg:: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_004.png
         :alt: fixed grid − whole, 8 random grids − whole
         :srcset: /auto_examples/06-learning/images/sphx_glr_03-networks-for-complex-volumes_004.png
         :class: sphx-glr-multi-img


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

 .. code-block:: none

    noisy                  relative error 0.1850
    whole volume           relative error 0.1147, departure from the whole volume 0.0000
    fixed grid             relative error 0.1169, departure from the whole volume 0.0321
    8 random grids         relative error 0.1219, departure from the whole volume 0.0487




.. GENERATED FROM PYTHON SOURCE LINES 360-377

The network, 0.14 million weights trained for a few minutes on patches of
one head, removes a third or more of the noise of the other head's volume
without blurring the white-matter tracts or the corpus callosum; a real
training set and a wider network remove more.

The departure of the fixed grid from the whole-volume result lies on the
planes between patches -- the seams -- and on the same planes at every call.
A shifted grid covers the volume with one more patch along each axis and so
has more boundaries, and a single call departs further from the
whole-volume result; but the boundaries move from call to call, and the
average of eight calls spreads the departure over the volume instead of
concentrating it on planes, where it would read as an anatomical edge.
Inside an iteration, which applies the denoiser once per step, one shifted
grid per call is enough: no plane receives the boundary error at every step.
The variance :func:`~bartorch.learning.moments` returns is a map of how much
the result depends on where the patches fall, one of the spreads of
:doc:`07-uncertainty`.

.. GENERATED FROM PYTHON SOURCE LINES 380-391

Size of the network
-------------------

The weights determine the storage footprint and, with the patch size, the
memory of a call. The default widths of :class:`~bartorch.learning.UNet`,
``(32, 64, 128, 256)``, give a three-dimensional network of a few million
weights. A series of frames -- a cine, a functional run -- is taken with
``frames=True``: the network then convolves each frame spatially and the
frames with a separate one-dimensional convolution, and never downsamples
the frame axis; ``periodic=True`` pads it circularly, which suits a cardiac
cycle. This factorization costs few weights beyond the spatial network.

.. GENERATED FROM PYTHON SOURCE LINES 392-404

.. code-block:: Python


    for name, candidate in (
        ("3D, widths (8, 16, 32)", net),
        ("3D, default widths", learning.UNet(6, spatial=3)),
        ("3D + frames, default widths", learning.UNet(6, spatial=3, frames=True, periodic=True)),
    ):
        count = sum(p.numel() for p in candidate.parameters())
        print(f"{name:<28} {count / 1e6:5.2f} M weights, {2 * count / 1e6:5.1f} MB in half precision")

    frames = torch.randn(1, 6, 10, 16, 16, 16)
    cine = learning.UNet(6, spatial=3, widths=(8, 16), frames=True, periodic=True)
    print("(n, channels, frames, z, y, x):", tuple(cine(frames).shape))




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

 .. code-block:: none

    3D, widths (8, 16, 32)        0.14 M weights,   0.3 MB in half precision
    3D, default widths            8.82 M weights,  17.6 MB in half precision
    3D + frames, default widths   9.48 M weights,  19.0 MB in half precision
    (n, channels, frames, z, y, x): (1, 6, 10, 16, 16, 16)





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

   **Total running time of the script:** (3 minutes 40.667 seconds)


.. _sphx_glr_download_auto_examples_06-learning_03-networks-for-complex-volumes.py:

.. only:: html

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

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

      :download:`Download Jupyter notebook: 03-networks-for-complex-volumes.ipynb <03-networks-for-complex-volumes.ipynb>`

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

      :download:`Download Python source code: 03-networks-for-complex-volumes.py <03-networks-for-complex-volumes.py>`

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

      :download:`Download zipped: 03-networks-for-complex-volumes.zip <03-networks-for-complex-volumes.zip>`


.. only:: html

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

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