torch_em.data.datasets.medical.reta

The RETA dataset contains annotations for retinal vessel segmentation in fundus images.

The dataset reuses the 81 fundus images from the first (segmentation) subset of the IDRiD dataset (see idrid.py) and adds pixel-level vessel masks obtained with a semi-automated coarse-to-fine annotation workflow. The public release also ships richer annotations (artery / vein masks, vascular skeletons, bifurcations, trees and abnormalities) as MATLAB '.mat' files for use with the authors' own "Computer Aided Retinal Labelling" (CARL) software; these graph / skeleton annotations are not covered by this module, which only exposes the ready-to-use binary vessel segmentation masks. Vessel masks for the 27-image test split are withheld by the authors for an online evaluation server, so only the 54-image training split has publicly available labels.

The dataset is located at https://doi.org/10.6084/m9.figshare.16960855 (CC BY 4.0). This dataset is from the publication https://doi.org/10.1038/s41597-022-01507-y. Please cite it if you use this dataset for your research.

  1"""The RETA dataset contains annotations for retinal vessel segmentation in fundus images.
  2
  3The dataset reuses the 81 fundus images from the first (segmentation) subset of the IDRiD dataset
  4(see `idrid.py`) and adds pixel-level vessel masks obtained with a semi-automated coarse-to-fine
  5annotation workflow. The public release also ships richer annotations (artery / vein masks,
  6vascular skeletons, bifurcations, trees and abnormalities) as MATLAB '.mat' files for use with the
  7authors' own "Computer Aided Retinal Labelling" (CARL) software; these graph / skeleton annotations
  8are not covered by this module, which only exposes the ready-to-use binary vessel segmentation
  9masks. Vessel masks for the 27-image test split are withheld by the authors for an online
 10evaluation server, so only the 54-image training split has publicly available labels.
 11
 12The dataset is located at https://doi.org/10.6084/m9.figshare.16960855 (CC BY 4.0).
 13This dataset is from the publication https://doi.org/10.1038/s41597-022-01507-y.
 14Please cite it if you use this dataset for your research.
 15"""
 16
 17import os
 18from glob import glob
 19from tqdm import tqdm
 20from pathlib import Path
 21from natsort import natsorted
 22from typing import Union, Literal, Tuple, List
 23
 24import numpy as np
 25import imageio.v3 as imageio
 26
 27from torch.utils.data import Dataset, DataLoader
 28
 29import torch_em
 30
 31from .. import util
 32
 33
 34URL = "https://ndownloader.figshare.com/files/31398340"
 35CHECKSUM = "02bd492a252d20c91c4f99f941b54160bd5db8a4b4b061ca272ab5697b818b4f"
 36
 37
 38def get_reta_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 39    """Download the RETA dataset.
 40
 41    Args:
 42        path: Filepath to a folder where the data is downloaded for further processing.
 43        download: Whether to download the data if it is not present.
 44
 45    Returns:
 46        Filepath where the data is downloaded.
 47    """
 48    data_dir = os.path.join(path, "images")
 49    if os.path.exists(data_dir):
 50        return data_dir
 51
 52    os.makedirs(path, exist_ok=True)
 53
 54    rar_path = os.path.join(path, "images.rar")
 55    util.download_source(path=rar_path, url=URL, download=download, checksum=CHECKSUM)
 56    util.unzip_rarfile(rar_path=rar_path, dst=path)
 57
 58    return data_dir
 59
 60
 61def get_reta_paths(
 62    path: Union[os.PathLike, str], split: Literal["train", "test"] = "train", download: bool = False,
 63) -> Tuple[List[str], List[str]]:
 64    """Get paths to the RETA data.
 65
 66    Args:
 67        path: Filepath to a folder where the data is downloaded for further processing.
 68        split: The choice of data split. The vessel masks for 'test' are withheld by the authors,
 69            so only 'train' has ground-truth to train / evaluate on.
 70        download: Whether to download the data if it is not present.
 71
 72    Returns:
 73        List of filepaths for the image data.
 74        List of filepaths for the label data.
 75    """
 76    data_dir = get_reta_data(path=path, download=download)
 77
 78    assert split in ["train", "test"], f"'{split}' is not a valid split."
 79    if split == "test":
 80        raise ValueError(
 81            "The vessel masks for the 'test' split are withheld by the RETA authors for online "
 82            "evaluation and are not publicly available. Please use the 'train' split instead."
 83        )
 84
 85    image_paths = natsorted(glob(os.path.join(data_dir, split, "img", "*.jpg")))
 86    raw_gt_paths = natsorted(glob(os.path.join(data_dir, split, "vessel", "*.png")))
 87    assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0
 88
 89    neu_gt_dir = os.path.join(data_dir, "..", "preprocessed", split)
 90    os.makedirs(neu_gt_dir, exist_ok=True)
 91
 92    gt_paths = []
 93    for image_path, raw_gt_path in tqdm(zip(image_paths, raw_gt_paths), total=len(image_paths), desc="Preprocessing labels"):  # noqa
 94        assert Path(image_path).stem == Path(raw_gt_path).stem.replace("_vessel", "")
 95
 96        gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif")
 97        gt_paths.append(gt_path)
 98        if os.path.exists(gt_path):
 99            continue
100
101        # the masks are stored as (near-)binary 3-channel pngs, i.e. non-zero pixels correspond to
102        # vessels in all channels alike. they are binarized into a uint8 (0, 1) single-channel map.
103        raw_gt = imageio.imread(raw_gt_path)
104        binary_gt = (raw_gt[..., 0] if raw_gt.ndim == 3 else raw_gt) > 0
105        imageio.imwrite(gt_path, binary_gt.astype(np.uint8))
106
107    return image_paths, gt_paths
108
109
110def get_reta_dataset(
111    path: Union[os.PathLike, str],
112    patch_shape: Tuple[int, int],
113    split: Literal["train", "test"] = "train",
114    resize_inputs: bool = False,
115    download: bool = False,
116    **kwargs
117) -> Dataset:
118    """Get the RETA dataset for retinal vessel segmentation in fundus images.
119
120    Args:
121        path: Filepath to a folder where the data is downloaded for further processing.
122        patch_shape: The patch shape to use for training.
123        split: The choice of data split. Only 'train' has publicly available vessel masks.
124        resize_inputs: Whether to resize the inputs to the expected patch shape.
125        download: Whether to download the data if it is not present.
126        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
127
128    Returns:
129        The segmentation dataset.
130    """
131    image_paths, gt_paths = get_reta_paths(path, split, download)
132
133    if resize_inputs:
134        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
135        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
136            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
137        )
138
139    return torch_em.default_segmentation_dataset(
140        raw_paths=image_paths,
141        raw_key=None,
142        label_paths=gt_paths,
143        label_key=None,
144        is_seg_dataset=False,
145        patch_shape=patch_shape,
146        **kwargs
147    )
148
149
150def get_reta_loader(
151    path: Union[os.PathLike, str],
152    batch_size: int,
153    patch_shape: Tuple[int, int],
154    split: Literal["train", "test"] = "train",
155    resize_inputs: bool = False,
156    download: bool = False,
157    **kwargs
158) -> DataLoader:
159    """Get the RETA dataloader for retinal vessel segmentation in fundus images.
160
161    Args:
162        path: Filepath to a folder where the data is downloaded for further processing.
163        batch_size: The batch size for training.
164        patch_shape: The patch shape to use for training.
165        split: The choice of data split. Only 'train' has publicly available vessel masks.
166        resize_inputs: Whether to resize the inputs to the expected patch shape.
167        download: Whether to download the data if it is not present.
168        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
169
170    Returns:
171        The DataLoader.
172    """
173    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
174    dataset = get_reta_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
175    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://ndownloader.figshare.com/files/31398340'
CHECKSUM = '02bd492a252d20c91c4f99f941b54160bd5db8a4b4b061ca272ab5697b818b4f'
def get_reta_data(path: Union[os.PathLike, str], download: bool = False) -> str:
39def get_reta_data(path: Union[os.PathLike, str], download: bool = False) -> str:
40    """Download the RETA dataset.
41
42    Args:
43        path: Filepath to a folder where the data is downloaded for further processing.
44        download: Whether to download the data if it is not present.
45
46    Returns:
47        Filepath where the data is downloaded.
48    """
49    data_dir = os.path.join(path, "images")
50    if os.path.exists(data_dir):
51        return data_dir
52
53    os.makedirs(path, exist_ok=True)
54
55    rar_path = os.path.join(path, "images.rar")
56    util.download_source(path=rar_path, url=URL, download=download, checksum=CHECKSUM)
57    util.unzip_rarfile(rar_path=rar_path, dst=path)
58
59    return data_dir

Download the RETA dataset.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • download: Whether to download the data if it is not present.
Returns:

Filepath where the data is downloaded.

def get_reta_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'] = 'train', download: bool = False) -> Tuple[List[str], List[str]]:
 62def get_reta_paths(
 63    path: Union[os.PathLike, str], split: Literal["train", "test"] = "train", download: bool = False,
 64) -> Tuple[List[str], List[str]]:
 65    """Get paths to the RETA data.
 66
 67    Args:
 68        path: Filepath to a folder where the data is downloaded for further processing.
 69        split: The choice of data split. The vessel masks for 'test' are withheld by the authors,
 70            so only 'train' has ground-truth to train / evaluate on.
 71        download: Whether to download the data if it is not present.
 72
 73    Returns:
 74        List of filepaths for the image data.
 75        List of filepaths for the label data.
 76    """
 77    data_dir = get_reta_data(path=path, download=download)
 78
 79    assert split in ["train", "test"], f"'{split}' is not a valid split."
 80    if split == "test":
 81        raise ValueError(
 82            "The vessel masks for the 'test' split are withheld by the RETA authors for online "
 83            "evaluation and are not publicly available. Please use the 'train' split instead."
 84        )
 85
 86    image_paths = natsorted(glob(os.path.join(data_dir, split, "img", "*.jpg")))
 87    raw_gt_paths = natsorted(glob(os.path.join(data_dir, split, "vessel", "*.png")))
 88    assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0
 89
 90    neu_gt_dir = os.path.join(data_dir, "..", "preprocessed", split)
 91    os.makedirs(neu_gt_dir, exist_ok=True)
 92
 93    gt_paths = []
 94    for image_path, raw_gt_path in tqdm(zip(image_paths, raw_gt_paths), total=len(image_paths), desc="Preprocessing labels"):  # noqa
 95        assert Path(image_path).stem == Path(raw_gt_path).stem.replace("_vessel", "")
 96
 97        gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif")
 98        gt_paths.append(gt_path)
 99        if os.path.exists(gt_path):
100            continue
101
102        # the masks are stored as (near-)binary 3-channel pngs, i.e. non-zero pixels correspond to
103        # vessels in all channels alike. they are binarized into a uint8 (0, 1) single-channel map.
104        raw_gt = imageio.imread(raw_gt_path)
105        binary_gt = (raw_gt[..., 0] if raw_gt.ndim == 3 else raw_gt) > 0
106        imageio.imwrite(gt_path, binary_gt.astype(np.uint8))
107
108    return image_paths, gt_paths

Get paths to the RETA data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split. The vessel masks for 'test' are withheld by the authors, so only 'train' has ground-truth to train / evaluate on.
  • download: Whether to download the data if it is not present.
Returns:

List of filepaths for the image data. List of filepaths for the label data.

def get_reta_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'] = 'train', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
111def get_reta_dataset(
112    path: Union[os.PathLike, str],
113    patch_shape: Tuple[int, int],
114    split: Literal["train", "test"] = "train",
115    resize_inputs: bool = False,
116    download: bool = False,
117    **kwargs
118) -> Dataset:
119    """Get the RETA dataset for retinal vessel segmentation in fundus images.
120
121    Args:
122        path: Filepath to a folder where the data is downloaded for further processing.
123        patch_shape: The patch shape to use for training.
124        split: The choice of data split. Only 'train' has publicly available vessel masks.
125        resize_inputs: Whether to resize the inputs to the expected patch shape.
126        download: Whether to download the data if it is not present.
127        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
128
129    Returns:
130        The segmentation dataset.
131    """
132    image_paths, gt_paths = get_reta_paths(path, split, download)
133
134    if resize_inputs:
135        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
136        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
137            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
138        )
139
140    return torch_em.default_segmentation_dataset(
141        raw_paths=image_paths,
142        raw_key=None,
143        label_paths=gt_paths,
144        label_key=None,
145        is_seg_dataset=False,
146        patch_shape=patch_shape,
147        **kwargs
148    )

Get the RETA dataset for retinal vessel segmentation in fundus images.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • split: The choice of data split. Only 'train' has publicly available vessel masks.
  • resize_inputs: Whether to resize the inputs to the expected patch shape.
  • download: Whether to download the data if it is not present.
  • kwargs: Additional keyword arguments for torch_em.default_segmentation_dataset.
Returns:

The segmentation dataset.

def get_reta_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'test'] = 'train', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
151def get_reta_loader(
152    path: Union[os.PathLike, str],
153    batch_size: int,
154    patch_shape: Tuple[int, int],
155    split: Literal["train", "test"] = "train",
156    resize_inputs: bool = False,
157    download: bool = False,
158    **kwargs
159) -> DataLoader:
160    """Get the RETA dataloader for retinal vessel segmentation in fundus images.
161
162    Args:
163        path: Filepath to a folder where the data is downloaded for further processing.
164        batch_size: The batch size for training.
165        patch_shape: The patch shape to use for training.
166        split: The choice of data split. Only 'train' has publicly available vessel masks.
167        resize_inputs: Whether to resize the inputs to the expected patch shape.
168        download: Whether to download the data if it is not present.
169        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
170
171    Returns:
172        The DataLoader.
173    """
174    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
175    dataset = get_reta_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
176    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the RETA dataloader for retinal vessel segmentation in fundus images.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • batch_size: The batch size for training.
  • patch_shape: The patch shape to use for training.
  • split: The choice of data split. Only 'train' has publicly available vessel masks.
  • resize_inputs: Whether to resize the inputs to the expected patch shape.
  • download: Whether to download the data if it is not present.
  • kwargs: Additional keyword arguments for torch_em.default_segmentation_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.