Learned reconstruction#

TL;DR

  • A network enters a reconstruction as the proximal step of an iteration: trained on its own and applied in a fixed iteration (plug-and-play), or trained through a fixed number of iterations (unrolled), or through the fixed point (deep equilibrium).

  • Complex multi-channel images are real channels to a network; the images stay on the host and a network runs on a device patch by patch, under mixed precision.

  • An unrolled network is trained in stages: the denoiser alone, then one iteration at a time, then end to end with checkpointing.

  • Without references, the acquired samples are split into a set to reconstruct from and a set to score on.

  • A spread from randomized reconstructions becomes an error bar once it is calibrated on references to a coverage.

A reconstruction solves \(\min_x \tfrac12 \|A x - y\|_2^2 + g(x)\) for an image \(x\) from data \(y\) acquired through an encoding \(A\) (Inverse problems and their solvers). A proximal iteration touches the regularizer \(g\) only through its proximal operator, a map from an image to a less noisy one, and a learned reconstruction replaces that map with a network. The encoding, the data consistency step and the iteration remain BART’s. This page states the ways the network is placed and trained, how a network is applied to data larger than GPU memory, and what an uncertainty estimate from it does and does not state.

Where the network enters#

Scheme

What is trained

On what

Iterations at inference

Plug-and-play

A denoiser, once

Pairs of a clean and a noisy image

As many as convergence needs

Unrolled

The denoiser and the step’s parameters

Data and references, through the iterations

The fixed number trained with

Deep equilibrium

As unrolled

As unrolled, differentiated at the fixed point

Until the fixed point

A plug-and-play denoiser is independent of the acquisition, so one network serves every protocol whose images resemble its training images, whatever the acceleration, sampling pattern or coil array; it is not adapted to the aliasing and g-factor noise a given undersampling produces. An unrolled network is trained on exactly those artefacts and needs few iterations, which is what bounds the reconstruction time on the scanner, and is tied to the encoding it was trained with. ImplicitPrior places a network in the proximal step of any block of bartorch.optim, Unrolled repeats a block a fixed number of times, and FixedPoint iterates it to its fixed point (Differentiation through reconstruction).

A step may depend on the iteration. In plug-and-play, the denoiser’s noise level \(\sigma_k\) is decreased over the iterations, so that early iterations remove the undersampling artefacts and later ones retain detail;[1] the ADMM penalty follows it as \(\rho_k = \lambda / \sigma_k^2\), the weight at which the proximal step of \(\lambda\,\phi\) is a Gaussian denoiser of variance \(\sigma_k^2\). A noise level or penalty given as a sequence is such a schedule. In an unrolled network, one network shared by the iterations is given the iteration index (steps=True of UNet) and modulates its features with it, which approaches the accuracy of a network per iteration at the storage of one. An iteration-dependent step has no fixed point, so FixedPoint refuses it.

Networks for complex, high-dimensional images#

A network operates on real tensors of shape (n, channels, *spatial). ComplexNet lays out the real and imaginary parts of every contrast, coefficient or channel of an image as its channels, so they are denoised jointly. Coefficient maps of a subspace basis differ in energy by orders of magnitude; whitening (normalize="whiten") decorrelates the channels and scales them to unit variance before the call and reverts both after it. A time series is taken with a frame axis (frames=True), convolved separately from the spatial axes and never downsampled.

A 3D volume, or a series of them such as a cine or a fingerprinting acquisition, rarely fits a network’s activations in the memory of a scanner’s GPU. The network is trained on patches, and Patchwise applies it to the whole image a few patches at a time: the image stays on the host, each group of patches is copied to the device, passed through the network under mixed precision (bfloat16 where the GPU supports it, float16 where it does not, such as a T4) and copied back. At inference the copies of one group run on a second stream while the network computes on another, so the transfers are hidden behind the computation. What the device holds is the network and two groups of patches. A convolutional network sees zeros past a patch boundary, which leaves seams at the boundaries of a fixed grid of patches; the grid is offset at random at every call, so successive iterations place the seams differently.

Training#

An unrolled network of many iterations records every iteration’s activations during the backward pass. Staged training bounds the memory:[2]

  1. The denoiser is trained alone, on pairs of an image and a degraded copy of it, with the iteration index it is to be used at.

  2. The unrolled network is trained one iteration at a time, each starting from the detached result of the previous one, with a loss on every iteration’s image weighted to increase geometrically.

  3. The whole network is trained end to end, with each iteration recomputed during the backward pass (checkpointing), so that memory holds one iteration’s activations.

Per-iteration training requires the block’s image to pass through that iteration’s denoiser, which holds for a proximal-gradient step and not for the image of an ADMM step, its x-update. Reconstruction runs each stage as a Lightning module; the items stay where the dataset put them, and a network inside decides where it runs.

Fully sampled references are rarely available for the data a learned reconstruction is most needed for: a dynamic or high-dimensional acquisition is undersampled because full sampling does not fit a breath-hold or a reasonable scan time. Self-supervision via data undersampling splits the acquired samples \(\Omega\) into disjoint sets \(\Theta\) and \(\Lambda\), reconstructs from \(\Theta\), and scores the reconstruction’s k-space on \(\Lambda\).[3] A new split is drawn at every step (split()), and the reconstruction at inference uses all of \(\Omega\).

Uncertainty#

Where the undersampling leaves the image underdetermined, a network fills in what its training data suggest, and a hallucinated structure is not distinguishable from anatomy in the image alone. A spread is obtained by repeating a randomized reconstruction and taking the voxel-wise variance (moments()): with dropout active in the network, from random subsets of the acquired samples, or with a random patch grid. Each spread measures one source of variability and none is the reconstruction error or a posterior distribution. Split conformal calibration relates it to the error: on held-out images with references, the factor \(q\) is the empirical quantile of \(|x - x_\mathrm{ref}| / s\) at the requested coverage (calibrate()), and \(q\,s\) is then an interval that contains the error at that rate on new images from the same distribution.[4] The rate holds on average over voxels and images, not at each voxel; how narrow the interval is where the reconstruction is reliable depends on how well the spread follows the error.

See also#

References#