learning.Patchwise#
- class bartorch.learning.Patchwise#
Bases:
ModuleA network applied to an image one patch at a time, on its own device.
The image stays where it is – normally the host – and is cut into non-overlapping patches over its trailing
len(patch)axes, which are copied to the device a few at a time, passed through the network under automatic mixed precision, and copied back into place. What the device holds at any moment is the network andbatchpatches, so a volume or a time series of volumes larger than the device’s memory goes through a network trained on patches of it. The result is on the input’s device, in single precision, and differentiable with respect to both the input and the weights.Before cutting, the image is padded with zeros by a random offset in front of each axis, drawn anew at every call from PyTorch’s generator, and by whatever completes the last patch behind it. The seams therefore fall in a different place at each iteration of a reconstruction and are not reinforced; with
shift=Falsethey are always at multiples of the patch. Drawing from PyTorch’s generator is what letsUnrolledrecompute a checkpointed step exactly.The network is moved to
devicebefore every call if it is not there, so a training loop that moves the whole model elsewhere does not move it; its parameters stay the same objects, and an optimizer holding them is unaffected.- Parameters:
net (nn.Module) – Called on
(n, channels, [frames,] *patch)asnet(x),net(x, sigma)or with the keywords this module is called with.patch (sequence of int) – Patch extent along each trailing axis, 2 or 3 of them.
device (torch.device or str, default=None) – Where the network runs; the first CUDA device when there is one, and the host otherwise.
dtype (torch.dtype or None or "auto", default="auto") – Precision of automatic mixed precision on the device.
"auto"is bfloat16 on a card that supports it and float16 on one that does not, such as a T4, and no mixed precision on the host.batch (int, default=1) – Patches per call of the network.
shift (bool, default=True) – Whether the patch grid is offset at random at every call.
overlap (bool, default=True) – Whether, without gradient and with the image on the host and the network on a card, the copies of one chunk of patches run while the network computes on another, on a second stream.
Examples
>>> net = learning.Patchwise(learning.UNet(10), patch=(64, 64, 64), batch=4) >>> prior = priors.ImplicitPrior(learning.ComplexNet(net, channels=1))
- forward()#
Apply the network over every patch of
x,(n, channels, ..., *patch-axes).sigmaand any keyword given one value per item ofxare repeated for each of that item’s patches.
- extra_repr()#
Return the extra representation of the module.
To print customized extra information, you should re-implement this method in your own modules. Both single-line and multi-line strings are acceptable.