priors.ImplicitPrior#

class bartorch.priors.ImplicitPrior#

Bases: Module

A denoiser substituted for a bartorch.priors regularizer.

The proximal operator is denoiser(x, sigma) independently of the step, as plug-and-play regularization defines it, or denoiser(x) where no sigma is given. Unlike a BART proximal operator this one is differentiable, so a solve containing it can be differentiated end to end.

The denoiser receives a leading batch axis, of length one for a single image. sigma is a torch.nn.Parameter, fixed until requires_grad_() is called.

A sequence of sigma is a schedule, one noise level per iteration and the last repeated beyond its end: the annealed plug-and-play iteration, which starts with a strong denoiser and weakens it as the data-consistent estimate improves. With step the denoiser is also given the iteration index, as denoiser(x, sigma, step=k) or denoiser(x, step=k), which is how a network shared by the iterations of an unrolled reconstruction adapts to each (bartorch.learning.UNet with steps=True). Either makes each iteration a different map, which bartorch.optim.FixedPoint refuses.

transform is the linear operator \(G\) of a term \(g(G x)\), as a regularizer built by BART also carries. The denoiser is then applied on the codomain of \(G\) rather than to the image, and the alternating-direction and primal-dual iterations split the variable at \(G x\), introducing one auxiliary variable and one dual variable per term and adding \(G^H G\) to the x-update. This accommodates a prior learned in a representation other than the optimization variable, such as contrast-weighted images obtained from subspace coefficient maps. Without it \(G\) is the identity and the denoiser is applied to the image.

Parameters:
  • denoiser (callable) – Called as denoiser(x), or as denoiser(x, sigma) when a sigma is given. Without spatial it receives the complex (batch, *shape) tensor as it stands, where shape is the image’s shape or, with a transform, the codomain of \(G\). With spatial it is an image-restoration network on real (n, channels, *spatial) planes of order unity – a deepinv, monai or local torch.nn.Module – and the conversion is done here.

  • sigma (float or sequence of float, default=None) – The noise level the denoiser is asked for, in the units its own convention states; a sequence gives one per iteration.

  • transform (LinearOperator, default=None) – \(G\), mapping the image to the domain the denoiser is applied on. Only the iterations given a term’s transform use it; see bartorch.priors.Regularizer.transform_is_identity().

  • spatial ({2, 3}, default=None) – Trailing axes the network operates on: 2 for a network trained on slices, 3 for one trained on volumes. The axes in front of them are folded into the network’s batch axis. None passes the complex tensor unconverted.

  • channels ({1, 2, 3}, default=1) – Input channels of the network: 1 for grayscale, 3 for RGB (a plane is repeated and the three outputs averaged), 2 for a network taking the real and imaginary planes jointly, as MoDL’s does. Only with spatial.

  • parts ({"channels", "separate", "magnitude"}, default=None) – How the complex values become real planes: the real and imaginary parts as the network’s two channels, each part denoised on its own in one call over a doubled batch, or the modulus denoised with the phase kept. Defaults to "channels" for two channels and "separate" otherwise. Only with spatial.

  • normalize (bool, default=True) – Scale each image to unit peak modulus around the call; sigma is then in units of that peak. Only with spatial.

  • step (bool, default=False) – Pass the iteration index to the denoiser as step.

Examples

>>> from deepinv.models import DRUNet
>>> prior = priors.ImplicitPrior(DRUNet(in_channels=1, out_channels=1), sigma=0.05, spatial=2)
>>> image = optim.fista(y, A, prior)
>>> priors.ImplicitPrior(denoiser, transform=linop.MultiplySum(basis, ...))
property stationary#

no schedule and no index.

Type:

Whether every iteration applies the same denoiser

prox()#
prox_shape()#
apply_transform()#
transform_is_identity()#
rewind()#

Examples using ImplicitPrior#

Plug-and-play denoisers

Plug-and-play denoisers

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