optim.FixedPoint#
- class bartorch.optim.FixedPoint#
Bases:
Moduleblockiterated to its fixed point: a deep-equilibrium model.The forward iterations run outside the graph until the moving state changes by at most
tolof its norm, or formax_itersteps. 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.iterationsis 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
Moduleinstance afterwards instead of this since the former takes care of running the registered hooks while the latter silently ignores them.