RandomPatchSampler#
- class deepinv.datasets.RandomPatchSampler(x_dir=None, y_dir=None, patch_size=32, file_format='.npy', ch_axis=None, dtype=torch.float32, loader=None, use_dict_output=False)[source]#
Bases:
ImageDatasetDataset for nD images that samples one random patch per image.
This dataset builds from one or two directories of nD images (must be of format
.npy,.nii(.gz), or.b2nd, ifloaderis not specified). On each epoch, it returns a randomly sampled patch of fixed size from each volume.Warning
This loader uses torch’s random functionality. To ensure reproducibility, set the DataLoader’s
generatorwith a fixed seed.Supported use cases:
Single-directory: provide only the ground-truth folder
x_diror measurement foldery_dir(returns patches from that directory).Paired-directory: provide both
x_dirandy_dir(returns matched patches from both).
Channel handling:
If
ch_axis=None: a singleton channel dimension is added at axis 0.If
ch_axis=0: images are assumed channel-first.If
ch_axis=-1: images are assumed channel-last and transposed to channel-first.Patches are never extracted along the channel axis (patch size for that axis is ignored).
Patch size handling:
Accepts either an integer (applied to all spatial dims) or a tuple.
If
patch_sizeis tuple, andpatch_size[i] == 1, this is equivalent to slicing across axis i (singleton at axis i will be squeezed). This can be used to e.g. extract 2D slices from a 3D volume.If tuple length is one less than the image ndim, the channel axis is auto-filled with
None.
Randomness & reproducibility:
Patch coordinates are drawn with Python’s
randommodule.To ensure deterministic behavior across workers, set the DataLoader’s
worker_init_fnorgeneratoraccording to the PyTorch reproducibility guidelines.
Notes:
All images must have the same dimensionality.
When both directories are provided, only files present in both are used.
Shapes of each file are checked for consistency (spatial not smaller than
patch_size+ channels remain consistent across files).
- Parameters:
x_dir (str, None) – Optional path to folder of ground-truth images. Required if
y_diris not given.y_dir (str, None) – Optional path to folder of measurement images. Required if
x_diris not given.patch_size (int, tuple) – Size of patches to extract. If int, applies the same size across all spatial dimensions (default
32)file_format (str) – File format to load. Other files are ignored (default
".npy").ch_axis (int, None) – Optional axis of the channel dimension. If None, a new singleton channel is added.
dtype (torch.dtype) – Data type to use when loading the images.
loader (Callable, None) – Optional custom loader function. Must accept path and the keyword
as_memmap, which will always be set to True. Must return an object that has shape attribute and returns anp.ndarraywhen sliced. If None, an internal loader is chosen based onfile_format.use_dict_output (bool) – whether to return output as dict with keys “x”, “y”, “params” instead of tuple (default
False).