learning.Reconstruction#

class bartorch.learning.Reconstruction#

Bases: LightningModule

A reconstruction network trained in one of the three stages of an unrolled network.

"denoiser" trains the network alone, on pairs of a degraded image and its reference, optionally with the iteration index and noise level it will meet in the unrolled network. "greedy" trains an Unrolled built with detach=True, taking a loss on each iteration’s image and back-propagating it before the next iteration runs, so memory holds one iteration; the losses are weighted geometrically, the last ratio times the first. "end-to-end" takes the loss on the last image alone, through the whole stack, which is affordable with checkpoint=True. Greedy training needs a block whose image passes through its own iteration’s denoiser, as ISTBlock’s gradient step followed by the denoiser does; ADMMBlock returns the image of its conjugate-gradient update, which depends on the previous iteration’s denoiser alone. Run in that order, each stage starting from the weights the previous one left, these are the staged training of Urman et al.

Items are dictionaries and a batch is a list of them (collate_fn=list), because operators and k-space vary between subjects:

  • "denoiser": input and target images, and optionally step, sigma and label passed to the network.

  • the unrolled stages: y and A, the measured data and its encoding operator; target, the reference image, or pattern for self-supervised training; optionally x0.

Every item may carry weights, multiplying the voxel-wise error of the default loss. Items stay where the dataset put them: the batch is not moved to the device the trainer runs on, since a Patchwise network or an operator whose operands are on the host decides itself what crosses to the card.

Without target the item is trained self-supervised: its pattern is split by split() at every step, the network reconstructs from pattern times the reconstruct part, and the loss is SSDU’s normalised l1-l2 loss on the held-out part of the k-space. The split drawn for validation is the same at every epoch.

Optimization is manual, so that a greedy step can release each iteration’s graph: gradients are accumulated over accumulate items, optionally clipped, and applied with Adam; the learning rate is reduced by factor after patience epochs without improvement of the logged val_loss, which lightning.pytorch.callbacks.EarlyStopping and lightning.pytorch.callbacks.ModelCheckpoint also monitor.

Parameters:
  • model (nn.Module) – The network for "denoiser", an Unrolled for the other stages.

  • stage ({"denoiser", "greedy", "end-to-end"}, default="end-to-end") – What is trained, and how.

  • loss (callable, default=None) – loss(prediction, target, item) for supervised items; the mean modulus of the error by default.

  • ratio (float, default=10.0) – Weight of the last iteration’s loss over the first’s, in the greedy stage.

  • lr (float, default=1e-3) – Adam’s learning rate.

  • weight_decay (float, default=0.0) – Adam’s weight decay.

  • factor (float, default=0.5) – Multiplier of the learning rate on a plateau of val_loss.

  • patience (int, default=10) – Epochs without improvement of val_loss before the learning rate is reduced.

  • accumulate (int, default=1) – Items whose gradients are summed before a step.

  • clip (float, default=None) – Maximum norm of the gradient.

  • fraction (float, default=0.4) – Share of the acquired samples held out, for self-supervised items.

  • split_options (dict, default=None) – Further keywords of split().

Examples

>>> stage = learning.Reconstruction(model, stage="greedy", accumulate=4)
>>> trainer = lightning.Trainer(max_epochs=50, callbacks=[EarlyStopping("val_loss")])
>>> trainer.fit(stage, DataLoader(train, collate_fn=list), DataLoader(valid, collate_fn=list))

References

Urman Y, Nishimura M, Abraham D, Cao X, Setsompop K. Fully 3D unrolled magnetic resonance fingerprinting reconstruction via staged pretraining and implicit gridding. arXiv:2601.17143, 2026.

transfer_batch_to_device()#

Leave the batch where the dataset put it.

configure_optimizers()#

Adam over the trainable parameters, reduced on a plateau of val_loss.

training_step()#

Back-propagate each item of a list and step every accumulate items.

validation_step()#

Log the mean loss of the items as val_loss.

on_validation_epoch_end()#

Step the plateau scheduler on val_loss.