priors.Regularizer#

class bartorch.priors.Regularizer#

Bases: ABC

Regularization functional with a proximal operator built by BART.

A regularizer represents \(g(G x)\): a linear transform \(G\) and the proximal operator of a functional \(g\) acting on its codomain. \(G\) is the identity for terms that penalize the image directly, and for terms carrying their own transform inside the proximal operator, such as Wavelet; it is a genuine operator for terms such as TotalVariation, whose functional acts on finite differences. Both come from BART’s opt_reg_configure.

Only the alternating-direction and primal-dual iterations are given \(G\); see transform_is_identity().

Both are built per image shape and cached, and handed to a solver as they are, so solving twice with the same term builds nothing the second time.

kind#

BART’s identifier for the regularizer, the letter opt_reg_configure selects it by.

Type:

str

axes#

Axes the term works over, as indices into the image’s shape.

Type:

tuple of int

joint_axes#

Axes along which the term acts jointly.

Type:

tuple of int

count#

Entries an NIHT term keeps; zero for every other term.

Type:

int

Examples

>>> term = priors.Wavelet((-1, -2), 0.01)
>>> z = term.prox(x, 0.5)                    # BART's proximal operator, step 0.5
>>> term.transform_is_identity((64, 64))     # G = I: every proximal solver takes it
True
>>> priors.TotalVariation((-1, -2), 0.01).transform_is_identity((64, 64))
False
build()#

The BART operator for this term over an image of C-order shape.

Built on first use for each shape. A term that draws random shifts is built once per item of a batch, so that each item draws the sequence it would draw alone. The returned handle is owned by this term and freed with it.

prox_shape()#

The C-order shape this term’s proximal operator works on.

The image’s, for a term that carries its own transform inside the proximal operator, a wavelet term included. A term whose transform is in front instead answers with the transform’s codomain: total variation thresholds the components of a gradient, and those are an axis of their own.

prox()#

prox_{gamma g}(x), the operator BART’s solvers apply.

Exposed directly so that an iteration written outside the library applies the same proximal operator the library would have.

Raises a ValueError for a tensor that requires a gradient: no backward pass is implemented for BART’s proximal operators, and treating one as the identity would zero the gradient path through the regularizer. Use ImplicitPrior for a differentiable proximal step, or detach() to hold this one fixed.

Parameters:
  • x (tensor) – Of prox_shape(), which for most terms is the image’s.

  • gamma (float, default=1.0) – The step the proximal operator is taken at. The term’s own weight is already in the operator, so this is only the step.

  • image_shape (tuple of int, default=None) – The image the term was configured for, when that is not what x is shaped like – which is the case for total variation, whose proximal operator works on the components of a gradient. By default x’s own shape, which is right for every other term.

  • item (int, default=0) – Which item of a batch x is; a term that draws random shifts keeps a generator per item.

Return type:

torch.Tensor

Notes

This is the proximal operator alone. A term may also carry a linear transform in front of it – see transform() – in which case a solver computes prox(transform(x)). BART’s iter2_ist applies the proximal operator to the image and ignores the transform entirely, so IST and FISTA admit only terms whose transform is the identity; a total-variation term is excluded, its proximal operator not being shaped like an image.

Examples

>>> priors.Wavelet((-1, -2), 0.01).prox(image, gamma=0.95)
apply_transform()#

This term’s transform applied to x, without making an operator of it.

transform() cannot answer for the gradient family – their components live on an axis past BART’s sixteen – and those are exactly the terms an alternating-direction solver is for. This applies the transform over the shapes prox_shape() reports instead, which a tensor can hold at any rank.

Parameters:
  • x (tensor) – The image for "forward", the proximal operator’s domain for "adjoint", the image for "normal".

  • image_shape (tuple of int, default=None) – What the term was configured for; by default x’s own shape, which is right whenever the transform starts from the image.

  • mode ({"forward", "adjoint", "normal"}, default='forward')

Return type:

torch.Tensor

Notes

Recorded for autograd when x carries a gradient: the backward pass is the transpose, which is the other of forward and adjoint, and normal itself. An unrolled network differentiates through the transform this way; the proximal operator behind it is BART’s and has no implemented backward pass, which is why ImplicitPrior exists.

transform()#

The operator BART puts in front of this term’s proximal operator.

The identity for a term that carries its transform inside the proximal operator instead. The shapes do not say which arrangement a term is: the Laplace term’s transform is a real convolution whose codomain is shaped like the image, so a caller inferring the arrangement from the shape would omit it without noticing.

Return type:

LinearOperator

rewind()#

Put this term’s own random generator back to where it started.

A wavelet threshold spins its transform by a shift drawn from a generator of its own, seeded at one when BART makes the operator. The tool builds a fresh operator per run; a term here is kept, so a solve rewinds it and a reused term answers as the tool does. Every solver does this before it iterates, whether the loop is BART’s or written out here; a term with no such generator is left alone.

transform_is_identity()#

Whether the term’s linear transform is the identity, as BART decides it.

Decided by linop_is_identity on the operator, not by comparing shapes: the Laplace term’s transform is a convolution on the image’s own shape, and total variation’s has a rank BART cannot hand over at all. iter2_chambolle_pock asks this of the first term before treating it as the primal proximal step rather than a dual one.

Return type:

bool

detach()#

This term with its proximal step detached, for a differentiated solve.

No backward pass is implemented for BART’s proximal operators, so prox() raises on a tensor that requires a gradient rather than contributing an incorrect one. The detached term thresholds exactly as before, and the gradient obtained is the one the iteration has with this term held fixed: the use is a solve in which a differentiable denoiser occupies one term and a BART term another, and only the denoiser is trained.

Examples

>>> optim.admm(y, A, [denoiser, priors.Wavelet((-1, -2), 0.01).detach()])