priors.ImplicitPrior#
- class bartorch.priors.ImplicitPrior#
Bases:
ModuleA denoiser substituted for a
bartorch.priorsregularizer.The proximal operator is
denoiser(x, sigma)independently of the step, as plug-and-play regularization defines it, ordenoiser(x)where nosigmais 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.
sigmais atorch.nn.Parameter, fixed untilrequires_grad_()is called.A sequence of
sigmais 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. Withstepthe denoiser is also given the iteration index, asdenoiser(x, sigma, step=k)ordenoiser(x, step=k), which is how a network shared by the iterations of an unrolled reconstruction adapts to each (bartorch.learning.UNetwithsteps=True). Either makes each iteration a different map, whichbartorch.optim.FixedPointrefuses.transformis 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 asdenoiser(x, sigma)when asigmais given. Withoutspatialit receives the complex(batch, *shape)tensor as it stands, whereshapeis the image’s shape or, with atransform, the codomain of \(G\). Withspatialit is an image-restoration network on real(n, channels, *spatial)planes of order unity – adeepinv,monaior localtorch.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.
Nonepasses 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 withspatial.normalize (bool, default=True) – Scale each image to unit peak modulus around the call;
sigmais then in units of that peak. Only withspatial.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()#