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

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

.. _sphx_glr_auto_examples_06-learning_07-uncertainty.py:


=======================
Uncertainty estimation
=======================

**Aim.** Attach a voxel-wise error bar to a learned reconstruction of
undersampled data, calibrated so that it contains the true error in a stated
fraction of voxels, and see where in the head the reconstruction is least
certain.

A learned reconstruction returns an image without saying where it may be
wrong. Where the undersampling leaves the image underdetermined, the network
fills in what its training data suggest, and a hallucinated structure looks
like real anatomy. An uncertainty map is obtained by randomizing the
reconstruction and repeating it: leaving dropout active in the network (Monte
Carlo dropout [#gal]_), reconstructing from random subsets of the acquired
phase encodes, or shifting a patch grid. Each spread measures one source of
variability, and none is the error itself. Split conformal calibration
[#angelopoulos]_ relates the spread to the error on held-out slices with
fully sampled references: it finds the factor by which the spread must be
multiplied for the interval to contain the error at a chosen rate, a
guarantee that holds whatever the spread measures.

**Learning objectives**

- Obtain a spread from Monte Carlo dropout and from k-space subsets with
  :func:`bartorch.learning.moments`.
- Calibrate it to a coverage with :func:`bartorch.learning.calibrate`, and
  check the coverage on other slices.
- Compare the calibrated interval with the error made.

It follows :doc:`06-annealed-plug-and-play`. This lesson ends the course; the
standalone examples of :doc:`../07-tours/index` apply the package to
individual problems.

.. GENERATED FROM PYTHON SOURCE LINES 37-146

.. 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 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 = 8
    DROPOUT = 0.1

    _ = torch.manual_seed(0)








.. GENERATED FROM PYTHON SOURCE LINES 147-153

Data
----

The slices, coils and fourfold undersampling of :doc:`04-staged-training`,
with sixteen slices of subject 4: the first eight to calibrate on, the last
eight to check the calibration on.

.. GENERATED FROM PYTHON SOURCE LINES 154-232

.. code-block:: Python


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

    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 = {
        part: [
            A(x) + NOISE * torch.randn(A.oshape, dtype=torch.complex64, generator=generator)
            for x in images
        ]
        for part, images in (("train", train_images), ("valid", valid_images))
    }








.. GENERATED FROM PYTHON SOURCE LINES 233-241

A network with dropout
----------------------

The iteration-conditioned unrolled network of :doc:`04-staged-training`,
with dropout in every residual block of its U-Net, trained end to end
against references. Dropout is a regularizer during training; left active
at inference it makes each reconstruction one draw from a family of
networks.

.. GENERATED FROM PYTHON SOURCE LINES 242-268

.. code-block:: Python


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

    items = {
        part: [{"y": y, "A": A, "target": x} for x, y in zip(images, kspace[part])]
        for part, images in (("train", train_images), ("valid", valid_images[:8]))
    }
    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(model, "end-to-end", lr=1e-3),
        DataLoader(items["train"], batch_size=4, shuffle=True, collate_fn=list),
        DataLoader(items["valid"], 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 269-279

Two spreads
-----------

:func:`~bartorch.learning.moments` calls a randomized reconstruction a
number of times and returns the voxel-wise mean and variance. The first
reconstruction leaves the dropout modules in training mode and everything
else in evaluation mode. The second keeps the network deterministic and
reconstructs each time from a random eighty per cent of the acquired lines,
drawn by :func:`~bartorch.learning.split`, which measures how much the image
depends on individual samples.

.. GENERATED FROM PYTHON SOURCE LINES 280-316

.. code-block:: Python


    model.eval()
    dropouts = [m for m in model.modules() if isinstance(m, torch.nn.Dropout)]
    acquired = lines.float()[:, None]


    def with_dropout(y):
        for module in dropouts:
            module.train()
        try:
            return model(y, A)
        finally:
            for module in dropouts:
                module.eval()


    def from_a_subset(y):
        keep, _ = learning.split(acquired, 0.2, keep=(8, 1))
        subset = linop.CartesianSense(
            sensitivities, (SIZE, SIZE), pattern=(pattern * keep).to(torch.complex64)
        )
        return model(y * keep, subset)


    spreads = {}
    for name, reconstruct in (("dropout", with_dropout), ("k-space subsets", from_a_subset)):
        means, deviations = [], []
        for y in kspace["valid"]:
            mean, variance = learning.moments(lambda: reconstruct(y), samples=8)
            means.append(mean)
            deviations.append(variance.sqrt())
        spreads[name] = (torch.stack(means), torch.stack(deviations))

    truth = torch.stack(valid_images)
    head = truth.abs() > 0.05








.. GENERATED FROM PYTHON SOURCE LINES 317-326

Calibration
-----------

On the first eight slices, whose references are known, the factor that makes
``|error| <= factor * spread`` hold for ninety per cent of the voxels in the
head is found with :func:`~bartorch.learning.calibrate`. On the other eight
the fraction of voxels whose error falls within ``factor * spread`` is
measured. Split conformal calibration guarantees that fraction on average
over voxels and subjects drawn alike, not voxel by voxel.

.. GENERATED FROM PYTHON SOURCE LINES 327-344

.. code-block:: Python


    COVERAGE = 0.9
    calibration, test = slice(0, 8), slice(8, 16)
    for name, (mean, deviation) in spreads.items():
        error = (mean - truth).abs()
        factor = learning.calibrate(
            error[calibration][head[calibration]], deviation[calibration][head[calibration]], COVERAGE
        )
        inside = error[test] <= factor * deviation[test]
        correlation = torch.corrcoef(
            torch.stack([error[test][head[test]], deviation[test][head[test]]])
        )[0, 1]
        print(
            f"{name:>16}: factor {factor:6.2f}, coverage {float(inside[head[test]].float().mean()):.3f}"
            f" (asked {COVERAGE}), correlation of error and spread {float(correlation):.2f}"
        )





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

 .. code-block:: none

             dropout: factor   8.57, coverage 0.906 (asked 0.9), correlation of error and spread 0.26
     k-space subsets: factor   3.82, coverage 0.897 (asked 0.9), correlation of error and spread 0.23




.. GENERATED FROM PYTHON SOURCE LINES 345-363

The coverage on the test slices is close to the one asked for, for both
spreads, although their factors differ: the calibration absorbs whatever
scale the spread has. What differs between them is how well the spread
follows the error voxel by voxel, which the correlation measures and the
maps below show. A spread that is large where the error is large gives
narrow intervals where the reconstruction is reliable and wide ones where it
is not; a spread unrelated to the error gives intervals of the right average
width in the wrong places.

The two spreads measure different things, and the maps show it. The
dropout interval is diffuse over the brain and follows neither its anatomy
nor the error. The k-space-subset interval is largest at the scalp and in
horizontal bands, the pattern of aliasing along the phase-encode direction
(vertical): removing lines moves the aliasing, and that is the variability
it records. The error itself is concentrated in the cortex. Both correlations
are weak, so the intervals are wider than the error over much of the white
matter and narrower than it in parts of the cortex. The coverage is met on
average over voxels, as the calibration guarantees, not voxel by voxel.

.. GENERATED FROM PYTHON SOURCE LINES 364-389




.. image-sg:: /auto_examples/06-learning/images/sphx_glr_07-uncertainty_001.png
   :alt: dropout mean, |error|, 90% interval, dropout, 90% interval, k-space subsets
   :srcset: /auto_examples/06-learning/images/sphx_glr_07-uncertainty_001.png
   :class: sphx-glr-single-img





.. GENERATED FROM PYTHON SOURCE LINES 390-400

References
----------

.. [#gal] Gal Y, Ghahramani Z. Dropout as a Bayesian approximation:
   representing model uncertainty in deep learning. *Proc Int Conf Mach
   Learn* 48:1050-1059 (2016). https://proceedings.mlr.press/v48/gal16.html

.. [#angelopoulos] Angelopoulos AN, Bates S. Conformal prediction: a gentle
   introduction. *Found Trends Mach Learn* 16(4):494-591 (2023).
   https://doi.org/10.1561/2200000101


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

   **Total running time of the script:** (1 minutes 11.981 seconds)


.. _sphx_glr_download_auto_examples_06-learning_07-uncertainty.py:

.. only:: html

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

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

      :download:`Download Jupyter notebook: 07-uncertainty.ipynb <07-uncertainty.ipynb>`

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

      :download:`Download Python source code: 07-uncertainty.py <07-uncertainty.py>`

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

      :download:`Download zipped: 07-uncertainty.zip <07-uncertainty.zip>`


.. only:: html

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

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