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)
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.
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.
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.