learning.Unrolled#

class bartorch.learning.Unrolled#

Bases: Module

Unrolled network formed by applying an iteration block iterations times.

The block implements one step of a BART iteration and this module implements the loop over it. With every parameter frozen the stack reproduces the corresponding solver bit for bit; calling requires_grad_() on a step size, a penalty weight or a denoiser’s weights makes the stack a trainable network. MoDL is this module applied to an ADMMBlock whose single term is a ImplicitPrior.

A single block shared by every iteration gives the weight sharing usual in unrolled networks; a sequence of blocks gives each iteration its own parameters.

detach and checkpoint select how the backward pass is taken. Neither alters the value the network computes:

  • end to end, the default, records the whole stack. Memory grows as the iteration count times the storage of one step.

  • detach=True starts each iteration from a detached state, so that the graph never spans two iterations. Combined with a loss on each image yielded by steps(), this is greedy per-iteration training, whose memory is independent of the iteration count.

  • checkpoint=True retains the states between iterations and recomputes the interior of a step during the backward pass. The gradient is the end-to-end one and each block is applied twice.

These correspond to the stages in which a large unrolled network is trained: a denoiser pretrained in isolation, then greedy per-iteration training, then end-to-end fine-tuning with gradient checkpointing. bartorch.optim.FixedPoint is an alternative with bounded memory: it drives the block to its fixed point and differentiates there, with no iteration count to unroll.

Parameters:
  • block (nn.Module or sequence of nn.Module) – One of bartorch.optim’s iteration blocks, or a sequence of them, one per iteration.

  • iterations (int, default=None) – Number of applications of a single shared block. Omitted for a sequence, whose length gives the count.

  • detach (bool, default=False) – Whether each iteration starts from a detached state.

  • checkpoint (bool, default=False) – Whether the interior of an iteration is recomputed during the backward pass rather than stored.

Notes

Checkpointing recomputes a step, and the gradient is correct only if the recomputation reproduces it. PyTorch’s random state is restored for the recomputation, so dropout in a denoiser is reproduced. A BART term that draws random shifts from BART’s own generator, such as Wavelet or LocallyLowRank with randshift=True, draws new shifts, and the gradient then does not correspond to the forward pass. This is not checked: checkpoint=True is valid only for steps that are deterministic or draw from PyTorch’s generator.

Examples

>>> block = optim.ADMMBlock(priors.ImplicitPrior(denoiser), rho=0.05, cg_maxiter=10)
>>> block.rho.requires_grad_()
>>> model = learning.Unrolled(block, iterations=10, checkpoint=True)
>>> model(kspace, A).abs().sub(target).square().mean().backward()
block()#

The block applied at iteration k: the shared one, or its own.

steps()#

Yield the image after each iteration in turn.

A loss taken on each yielded image, with detach set, gives greedy per-iteration training. forward() returns the last of them.

forward()#

Reconstruct y, returning the image left by the last iteration.

Parameters:
  • y (torch.Tensor) – Measured data, with or without a leading batch axis.

  • A (LinearOperator) – Encoding operator, shared by every item of a batch.

  • x0 (torch.Tensor, default=None) – Starting point of the iteration; zero by default, as in BART.

Examples using Unrolled#

MoDL, on BART’s ADMM

MoDL, on BART's ADMM

Staged training of an unrolled network

Staged training of an unrolled network

Training without a reference

Training without a reference

Annealed plug-and-play

Annealed plug-and-play

Uncertainty estimation

Uncertainty estimation