learning.Unrolled#
- class bartorch.learning.Unrolled#
Bases:
ModuleUnrolled network formed by applying an iteration block
iterationstimes.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 anADMMBlockwhose single term is aImplicitPrior.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.
detachandcheckpointselect 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=Truestarts each iteration from a detached state, so that the graph never spans two iterations. Combined with a loss on each image yielded bysteps(), this is greedy per-iteration training, whose memory is independent of the iteration count.checkpoint=Trueretains 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.FixedPointis 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
WaveletorLocallyLowRankwithrandshift=True, draws new shifts, and the gradient then does not correspond to the forward pass. This is not checked:checkpoint=Trueis 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
detachset, 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.