learning.ComplexNet

learning.ComplexNet#

class bartorch.learning.ComplexNet#

Bases: Module

A network on real channels applied to complex images whose leading axes are channels.

The input is (batch, *channels, *spatial), or (batch, *channels, frames, *spatial) with frames. The complex values and the channels axes become the network’s 2 * prod(channels) real channels, real parts first, in the order of as_real() applied to each item; with frames the frame axis is kept as the axis in front of the spatial ones, which is what UNet with frames=True takes. This is the layout for subspace coefficient maps, which are denoised jointly as one multi-channel image rather than one coefficient at a time.

normalize brings each item to order unity around the call, with statistics taken without gradient:

  • "peak" divides by the peak modulus, so that sigma is in units of that peak.

  • "whiten" subtracts each real channel’s mean and multiplies by the inverse square root of the channels’ covariance, so that the channels the network sees are uncorrelated and of unit variance. Coefficient maps of a subspace basis differ in energy by orders of magnitude, and this balances them. sigma then has no fixed unit.

Parameters:
  • net (callable) – Called on real (batch, 2 * prod(channels), [frames,] *spatial) tensors as net(x), net(x, sigma), and with any further keyword this module is called with, such as step.

  • spatial ({2, 3}, default=3) – Number of trailing spatial axes.

  • channels (int, default=0) – Number of axes in front of the spatial ones (and of the frame axis) folded into the channels.

  • frames (bool, default=False) – Whether the axis in front of the spatial ones is a frame axis.

  • normalize ({"peak", "whiten", None}, default="peak") – How each item is scaled around the call.

Examples

>>> net = learning.ComplexNet(learning.UNet(10, steps=True), channels=1, normalize="whiten")
>>> prior = priors.ImplicitPrior(net, step=True)      # (5, z, y, x) coefficient maps
forward()#

Apply the network, returning input’s shape and dtype.

A real input is taken as complex with a zero imaginary part, and the real part of the result is returned.