torch_em.data.datasets.medical.fives

The FIVES dataset contains annotations for retinal vessel segmentation in high-resolution fundus images across four categories: normal, age-related macular degeneration (AMD), diabetic retinopathy (DR) and glaucoma.

This dataset is from the publication https://doi.org/10.1038/s41597-022-01564-3. The dataset is hosted on figshare at https://doi.org/10.6084/m9.figshare.19688169.v1 and is licensed under CC BY 4.0. Please cite the publication above if you use this dataset for your research.

  1"""The FIVES dataset contains annotations for retinal vessel segmentation in high-resolution
  2fundus images across four categories: normal, age-related macular degeneration (AMD), diabetic
  3retinopathy (DR) and glaucoma.
  4
  5This dataset is from the publication https://doi.org/10.1038/s41597-022-01564-3.
  6The dataset is hosted on figshare at https://doi.org/10.6084/m9.figshare.19688169.v1 and is
  7licensed under CC BY 4.0. Please cite the publication above if you use this dataset for your
  8research.
  9"""
 10
 11import os
 12from glob import glob
 13from pathlib import Path
 14from typing import Union, Tuple, Literal, List
 15
 16import imageio.v3 as imageio
 17
 18from torch.utils.data import Dataset, DataLoader
 19
 20import torch_em
 21
 22from .. import util
 23
 24
 25URL = "https://ndownloader.figshare.com/files/34969398"
 26CHECKSUM = "be72f9af286b107bcebcc08a9dae7fc55c3fb0959409b689e14c72f9fdc4ad8e"
 27
 28CATEGORIES = {"N": "normal", "A": "amd", "D": "dr", "G": "glaucoma"}
 29
 30
 31def get_fives_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 32    """Download the FIVES dataset.
 33
 34    Args:
 35        path: Filepath to a folder where the data is downloaded for further processing.
 36        download: Whether to download the data if it is not present.
 37
 38    Returns:
 39        Filepath where the data is downloaded.
 40    """
 41    data_dir = os.path.join(path, "FIVES A Fundus Image Dataset for AI-based Vessel Segmentation")
 42    if os.path.exists(data_dir):
 43        return data_dir
 44
 45    os.makedirs(path, exist_ok=True)
 46
 47    rar_path = os.path.join(path, "fives.rar")
 48    util.download_source(path=rar_path, url=URL, download=download, checksum=CHECKSUM)
 49    util.unzip_rarfile(rar_path=rar_path, dst=path)
 50
 51    return data_dir
 52
 53
 54def _get_fives_ground_truth(data_dir, split):
 55    gt_paths = sorted(glob(os.path.join(data_dir, split, "Ground truth", "*.png")))
 56
 57    neu_gt_dir = os.path.join(data_dir, split, "gt")
 58    if os.path.exists(neu_gt_dir):
 59        return sorted(glob(os.path.join(neu_gt_dir, "*.tif")))
 60    else:
 61        os.makedirs(neu_gt_dir, exist_ok=True)
 62
 63    neu_gt_paths = []
 64    for gt_path in gt_paths:
 65        gt = imageio.imread(gt_path)
 66        if gt.ndim == 3:
 67            gt = gt[..., 0]
 68        neu_gt_path = os.path.join(neu_gt_dir, Path(os.path.split(gt_path)[-1]).with_suffix(".tif"))
 69        imageio.imwrite(neu_gt_path, (gt > 0).astype("uint8"))
 70        neu_gt_paths.append(neu_gt_path)
 71
 72    return sorted(neu_gt_paths)
 73
 74
 75def get_fives_paths(
 76    path: Union[os.PathLike, str],
 77    split: Literal["train", "test"],
 78    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
 79    download: bool = False,
 80) -> Tuple[List[str], List[str]]:
 81    """Get paths to the FIVES data.
 82
 83    Args:
 84        path: Filepath to a folder where the data is downloaded for further processing.
 85        split: The choice of data split. Either 'train' or 'test'.
 86        category: The choice of disease category. One of 'normal', 'amd', 'dr', 'glaucoma' or 'all'.
 87        download: Whether to download the data if it is not present.
 88
 89    Returns:
 90        List of filepaths for the image data.
 91        List of filepaths for the label data.
 92    """
 93    if split not in ("train", "test"):
 94        raise ValueError(f"'{split}' is not a valid split.")
 95
 96    if category == "all":
 97        prefixes = list(CATEGORIES)
 98    else:
 99        matches = [k for k, v in CATEGORIES.items() if v == category]
100        if not matches:
101            valid = list(CATEGORIES.values()) + ["all"]
102            raise ValueError(f"'{category}' is not a valid category. Choose from {valid}.")
103        prefixes = matches
104
105    data_dir = get_fives_data(path=path, download=download)
106
107    image_paths = sorted(glob(os.path.join(data_dir, split, "Original", "*.png")))
108    gt_paths = _get_fives_ground_truth(data_dir, split)
109
110    image_paths = [p for p in image_paths if os.path.splitext(os.path.basename(p))[0].split("_")[-1] in prefixes]
111    gt_paths = [p for p in gt_paths if os.path.splitext(os.path.basename(p))[0].split("_")[-1] in prefixes]
112
113    assert len(image_paths) == len(gt_paths) and len(image_paths) > 0
114
115    return image_paths, gt_paths
116
117
118def get_fives_dataset(
119    path: Union[os.PathLike, str],
120    patch_shape: Tuple[int, int],
121    split: Literal["train", "test"],
122    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
123    resize_inputs: bool = False,
124    download: bool = False,
125    **kwargs
126) -> Dataset:
127    """Get the FIVES dataset for segmentation of retinal blood vessels in high-resolution fundus images.
128
129    Args:
130        path: Filepath to a folder where the data is downloaded for further processing.
131        patch_shape: The patch shape to use for training.
132        split: The choice of data split. Either 'train' or 'test'.
133        category: The choice of disease category.
134        resize_inputs: Whether to resize the inputs to the expected patch shape.
135        download: Whether to download the data if it is not present.
136        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
137
138    Returns:
139        The segmentation dataset.
140    """
141    image_paths, gt_paths = get_fives_paths(path=path, split=split, category=category, download=download)
142
143    if resize_inputs:
144        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
145        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
146            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
147        )
148
149    return torch_em.default_segmentation_dataset(
150        raw_paths=image_paths,
151        raw_key=None,
152        label_paths=gt_paths,
153        label_key=None,
154        patch_shape=patch_shape,
155        is_seg_dataset=False,
156        **kwargs
157    )
158
159
160def get_fives_loader(
161    path: Union[os.PathLike, str],
162    batch_size: int,
163    patch_shape: Tuple[int, int],
164    split: Literal["train", "test"],
165    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
166    resize_inputs: bool = False,
167    download: bool = False,
168    **kwargs
169) -> DataLoader:
170    """Get the FIVES dataloader for segmentation of retinal blood vessels in high-resolution fundus images.
171
172    Args:
173        path: Filepath to a folder where the data is downloaded for further processing.
174        batch_size: The batch size for training.
175        patch_shape: The patch shape to use for training.
176        split: The choice of data split. Either 'train' or 'test'.
177        category: The choice of disease category.
178        resize_inputs: Whether to resize the inputs to the expected patch shape.
179        download: Whether to download the data if it is not present.
180        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
181
182    Returns:
183        The DataLoader.
184    """
185    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
186    dataset = get_fives_dataset(path, patch_shape, split, category, resize_inputs, download, **ds_kwargs)
187    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://ndownloader.figshare.com/files/34969398'
CHECKSUM = 'be72f9af286b107bcebcc08a9dae7fc55c3fb0959409b689e14c72f9fdc4ad8e'
CATEGORIES = {'N': 'normal', 'A': 'amd', 'D': 'dr', 'G': 'glaucoma'}
def get_fives_data(path: Union[os.PathLike, str], download: bool = False) -> str:
32def get_fives_data(path: Union[os.PathLike, str], download: bool = False) -> str:
33    """Download the FIVES dataset.
34
35    Args:
36        path: Filepath to a folder where the data is downloaded for further processing.
37        download: Whether to download the data if it is not present.
38
39    Returns:
40        Filepath where the data is downloaded.
41    """
42    data_dir = os.path.join(path, "FIVES A Fundus Image Dataset for AI-based Vessel Segmentation")
43    if os.path.exists(data_dir):
44        return data_dir
45
46    os.makedirs(path, exist_ok=True)
47
48    rar_path = os.path.join(path, "fives.rar")
49    util.download_source(path=rar_path, url=URL, download=download, checksum=CHECKSUM)
50    util.unzip_rarfile(rar_path=rar_path, dst=path)
51
52    return data_dir

Download the FIVES 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_fives_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'], category: Literal['normal', 'amd', 'dr', 'glaucoma', 'all'] = 'all', download: bool = False) -> Tuple[List[str], List[str]]:
 76def get_fives_paths(
 77    path: Union[os.PathLike, str],
 78    split: Literal["train", "test"],
 79    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
 80    download: bool = False,
 81) -> Tuple[List[str], List[str]]:
 82    """Get paths to the FIVES data.
 83
 84    Args:
 85        path: Filepath to a folder where the data is downloaded for further processing.
 86        split: The choice of data split. Either 'train' or 'test'.
 87        category: The choice of disease category. One of 'normal', 'amd', 'dr', 'glaucoma' or 'all'.
 88        download: Whether to download the data if it is not present.
 89
 90    Returns:
 91        List of filepaths for the image data.
 92        List of filepaths for the label data.
 93    """
 94    if split not in ("train", "test"):
 95        raise ValueError(f"'{split}' is not a valid split.")
 96
 97    if category == "all":
 98        prefixes = list(CATEGORIES)
 99    else:
100        matches = [k for k, v in CATEGORIES.items() if v == category]
101        if not matches:
102            valid = list(CATEGORIES.values()) + ["all"]
103            raise ValueError(f"'{category}' is not a valid category. Choose from {valid}.")
104        prefixes = matches
105
106    data_dir = get_fives_data(path=path, download=download)
107
108    image_paths = sorted(glob(os.path.join(data_dir, split, "Original", "*.png")))
109    gt_paths = _get_fives_ground_truth(data_dir, split)
110
111    image_paths = [p for p in image_paths if os.path.splitext(os.path.basename(p))[0].split("_")[-1] in prefixes]
112    gt_paths = [p for p in gt_paths if os.path.splitext(os.path.basename(p))[0].split("_")[-1] in prefixes]
113
114    assert len(image_paths) == len(gt_paths) and len(image_paths) > 0
115
116    return image_paths, gt_paths

Get paths to the FIVES data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split. Either 'train' or 'test'.
  • category: The choice of disease category. One of 'normal', 'amd', 'dr', 'glaucoma' or 'all'.
  • 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_fives_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'], category: Literal['normal', 'amd', 'dr', 'glaucoma', 'all'] = 'all', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
119def get_fives_dataset(
120    path: Union[os.PathLike, str],
121    patch_shape: Tuple[int, int],
122    split: Literal["train", "test"],
123    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
124    resize_inputs: bool = False,
125    download: bool = False,
126    **kwargs
127) -> Dataset:
128    """Get the FIVES dataset for segmentation of retinal blood vessels in high-resolution fundus images.
129
130    Args:
131        path: Filepath to a folder where the data is downloaded for further processing.
132        patch_shape: The patch shape to use for training.
133        split: The choice of data split. Either 'train' or 'test'.
134        category: The choice of disease category.
135        resize_inputs: Whether to resize the inputs to the expected patch shape.
136        download: Whether to download the data if it is not present.
137        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
138
139    Returns:
140        The segmentation dataset.
141    """
142    image_paths, gt_paths = get_fives_paths(path=path, split=split, category=category, download=download)
143
144    if resize_inputs:
145        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
146        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
147            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
148        )
149
150    return torch_em.default_segmentation_dataset(
151        raw_paths=image_paths,
152        raw_key=None,
153        label_paths=gt_paths,
154        label_key=None,
155        patch_shape=patch_shape,
156        is_seg_dataset=False,
157        **kwargs
158    )

Get the FIVES dataset for segmentation of retinal blood vessels in high-resolution 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. Either 'train' or 'test'.
  • category: The choice of disease category.
  • 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_fives_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'test'], category: Literal['normal', 'amd', 'dr', 'glaucoma', 'all'] = 'all', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
161def get_fives_loader(
162    path: Union[os.PathLike, str],
163    batch_size: int,
164    patch_shape: Tuple[int, int],
165    split: Literal["train", "test"],
166    category: Literal["normal", "amd", "dr", "glaucoma", "all"] = "all",
167    resize_inputs: bool = False,
168    download: bool = False,
169    **kwargs
170) -> DataLoader:
171    """Get the FIVES dataloader for segmentation of retinal blood vessels in high-resolution fundus images.
172
173    Args:
174        path: Filepath to a folder where the data is downloaded for further processing.
175        batch_size: The batch size for training.
176        patch_shape: The patch shape to use for training.
177        split: The choice of data split. Either 'train' or 'test'.
178        category: The choice of disease category.
179        resize_inputs: Whether to resize the inputs to the expected patch shape.
180        download: Whether to download the data if it is not present.
181        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
182
183    Returns:
184        The DataLoader.
185    """
186    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
187    dataset = get_fives_dataset(path, patch_shape, split, category, resize_inputs, download, **ds_kwargs)
188    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the FIVES dataloader for segmentation of retinal blood vessels in high-resolution 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. Either 'train' or 'test'.
  • category: The choice of disease category.
  • 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.