learning.Reconstruction#
- class bartorch.learning.Reconstruction#
Bases:
LightningModuleA 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 anUnrolledbuilt withdetach=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 lastratiotimes the first."end-to-end"takes the loss on the last image alone, through the whole stack, which is affordable withcheckpoint=True. Greedy training needs a block whose image passes through its own iteration’s denoiser, asISTBlock’s gradient step followed by the denoiser does;ADMMBlockreturns 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":inputandtargetimages, and optionallystep,sigmaandlabelpassed to the network.the unrolled stages:
yandA, the measured data and its encoding operator;target, the reference image, orpatternfor self-supervised training; optionallyx0.
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 aPatchwisenetwork or an operator whose operands are on the host decides itself what crosses to the card.Without
targetthe item is trained self-supervised: itspatternis split bysplit()at every step, the network reconstructs frompatterntimes 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
accumulateitems, optionally clipped, and applied with Adam; the learning rate is reduced byfactorafterpatienceepochs without improvement of the loggedval_loss, whichlightning.pytorch.callbacks.EarlyStoppingandlightning.pytorch.callbacks.ModelCheckpointalso monitor.- Parameters:
model (nn.Module) – The network for
"denoiser", anUnrolledfor 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_lossbefore 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
accumulateitems.
- validation_step()#
Log the mean loss of the items as
val_loss.
- on_validation_epoch_end()#
Step the plateau scheduler on
val_loss.