nlop.IRGNMBlock#

class bartorch.nlop.IRGNMBlock#

Bases: Module

One step of iteratively regularized Gauss-Newton, noir’s Gauss-Newton step.

Without inner the step is BART’s first form,

x = xn + (DF^H DF + alpha)^-1 [DF^H (y - F(xn)) - alpha (xn - xref)],

with the inverse computed by conjugate gradients inside the library and differentiated implicitly. With inner it is the second form: the solver minimizes ||DF u - r||^2 + alpha ||u||^2 + R(u) for DF linearized at xn and r = y - F(xn) + DF (xn - xref), and x = u + xref. The weight then decays as alpha <- (alpha - alpha_min) / redu + alpha_min, and in the second form is kept above alpha_min0.

state = block.start(y, F, x0, xref) prepares the data and the model, state = block(state, F) takes one step, and block.output(state, F) returns the unknowns. alpha is the weight of the first step: the block taking step k applies it decayed k times, so one block looped, or a stack of blocks with the same alpha, is BART’s schedule, and a stack whose blocks learn alpha learns a weight per step. F needs a _bundled. A leading batch axis on y is a batch of independent items, each stepped on its own; a model that holds its items, as CoilSense(..., items=True) does, is stepped as one, with each item’s inner problem solved on its own.

Parameters:
  • alpha (float, default=1.0) – Initial Tikhonov weight.

  • alpha_min (float, default=0.0) – What the weight decays towards.

  • alpha_min0 (float, default=0.0) – Floor the decayed weight is never taken below; second form only.

  • redu (float, default=2.0) – Factor the weight is divided by after each step.

  • cg_maxiter (int, default=30) – The first form’s conjugate-gradient iterations, iter_conjgrad_conf’s maxiter.

  • cg_tol (float, default=0.0) – Its tolerance, iter_conjgrad_conf’s tol. A nonzero tol is refused by BART once the step is differentiated.

  • cg_lambda (float, default=0.0) – Its Tikhonov weight, iter_conjgrad_conf’s l2lambda.

  • inner (solver, default=None) – A configured solver from bartorch.optim for the linearized problem; its regularizers are R.

  • fuse (bool, default=True) – Lower a coil composition so that its encoding is applied once as its normal operator; see plan().

Notes

alpha and redu are torch.nn.Parameter objects, frozen until requires_grad_(). The first form is differentiable by the data, the iterate, the centre and alpha. The second form is differentiable by the data, the iterate, the centre and the inner solver’s own settings and priors; it gives alpha to the solver as a number, so alpha is held fixed there.

plan()#

What F was lowered into, as a Plan.

plan.derivative says where the derivative came from, plan.domain whether the step works against the normal operator or applies the encoding forward and adjoint, plan.encoding the linear part’s own plan, and plan.fused whether the coil model was rewritten.

start()#

The run’s state: the prepared data, the start, the centre and the weight.

x0 and xref are a state (the unknowns laid end to end), the one unknown of a single-input model at its shape, or a tuple of unknowns. Without x0 the start is nlinv’s: the first unknown ones, the rest zero. Without xref the steps are regularized towards zero.

forward()#

Define the computation performed at every call.

Should be overridden by all subclasses.

Note

Although the recipe for forward pass needs to be defined within this function, one should call the Module instance afterwards instead of this since the former takes care of running the registered hooks while the latter silently ignores them.

output()#

The unknowns: a tensor for a model of one input, else one per input.