optim.FixedPoint

optim.FixedPoint#

class bartorch.optim.FixedPoint#

Bases: Module

block iterated to its fixed point: a deep-equilibrium model.

The forward iterations run outside the graph until the moving state changes by at most tol of its norm, or for max_iter steps. Differentiation is implicit rather than through the unrolled run: one step is taken from the fixed point, and the vector-Jacobian products of that step are iterated to their own fixed point, so memory does not grow with the iteration count. iterations is the last forward pass’s count.

A step that is not the same map at every iteration – FISTA’s momentum, a moving rho, adaptive or decaying step sizes – has no fixed point and is rejected.

Examples

>>> deq = optim.FixedPoint(optim.PRIDUBlock(priors.ImplicitPrior(net, 0.05)), tol=1e-4)
>>> deq(kspace, A).abs().sub(target).square().sum().backward()
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.