learning.UNet

Contents

learning.UNet#

class bartorch.learning.UNet#

Bases: Module

Residual UNet on real (n, channels, *spatial) tensors, for denoising and restoration.

Each level holds blocks residual blocks of two 3-wide convolutions with group normalization and SiLU; levels are joined by strided convolutions on the way down and by nearest-neighbour upsampling and a convolution on the way up. With residual the network returns its input plus a correction, and the last convolution starts at zero, so that an untrained network is the identity and an unrolled iteration around it starts as the plain iteration.

With frames the input is (n, channels, frames, *spatial) and every convolution is factorised into a spatial one per frame and a temporal one per voxel; the frame axis is never downsampled. This is the network for a time series or any other axis that is not spatial but is correlated from one entry to the next.

Conditioning on the iteration index (steps), the noise level (noise) and a class label (classes) is embedded, summed and applied in every block as a feature-wise affine modulation (FiLM), (1 + gamma) h + beta. A network shared by all the iterations of an unrolled reconstruction is conditioned on the index; a denoiser applied over an annealed schedule, on the noise level; one serving several contrasts, on the contrast.

Parameters:
  • in_channels (int) – Input channels: 2 for one complex image as real and imaginary parts, 2 k for k complex images such as subspace coefficients.

  • out_channels (int, default=None) – Output channels; in_channels when omitted.

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

  • widths (sequence of int, default=(32, 64, 128, 256)) – Channels at each level, from full resolution down. Every spatial extent is padded to a multiple of 2 ** (len(widths) - 1).

  • blocks (int, default=1) – Residual blocks per level on each side.

  • groups (int, default=8) – Groups of the group normalization.

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

  • periodic (bool, default=False) – Whether the frame axis wraps around, as a cardiac cycle does.

  • steps (bool, default=False) – Whether the network takes the iteration index, step.

  • noise (bool, default=False) – Whether the network takes the noise level, sigma.

  • classes (int, default=None) – Number of class labels the network takes, as label.

  • features (int, default=64) – Width of the conditioning embedding.

  • residual (bool, default=True) – Whether the output is the input plus a correction. Needs out_channels == in_channels.

  • dropout (float, default=0.0) – Dropout probability inside each residual block. Left active at inference, it makes repeated reconstructions differ, which moments() turns into a spread.

Examples

>>> net = learning.UNet(10, steps=True)              # five complex coefficients, 3D
>>> net(planes, step=2).shape == planes.shape
True
>>> cine = learning.UNet(2, spatial=2, frames=True, periodic=True)
forward()#

Apply the network.

Parameters:
  • x (torch.Tensor) – (n, channels, *spatial), or (n, channels, frames, *spatial) with frames.

  • sigma (float or torch.Tensor, default=None) – Noise level, one or one per item; only with noise.

  • step (int or torch.Tensor, default=None) – Iteration index, one or one per item; only with steps.

  • label (int or torch.Tensor, default=None) – Class label, one or one per item; only with classes.