torch_em.data.datasets.medical.deeprt
The DeepRT dataset contains annotations for retinal tissue segmentation in optical coherence tomography (OCT) B-scans, originally released to support self-supervised pretraining for diabetic retinopathy classification.
The dataset ships 1,286 OCT images (acquired with Spectralis and Topcon devices) with binary retinal tissue masks, pre-split into 'train', 'validation' and 'test' (744 / 265 / 277 images respectively) via the 'file_names_complete' mapping files shipped in the archive. NOTE: the associated publication reports 1,009 "semantic segmentations", a smaller curated subset; this loader exposes the full 1,286 labeled images shipped in the 'thickness_segmentation_data' archive.
The data is located at https://doi.org/10.5281/zenodo.3626020, released under a CC-BY-4.0 license.
This dataset is from the publication https://doi.org/10.1038/s42256-020-00247-1. Please cite it if you use this dataset for your research.
1"""The DeepRT dataset contains annotations for retinal tissue segmentation in optical coherence 2tomography (OCT) B-scans, originally released to support self-supervised pretraining for diabetic 3retinopathy classification. 4 5The dataset ships 1,286 OCT images (acquired with Spectralis and Topcon devices) with binary retinal 6tissue masks, pre-split into 'train', 'validation' and 'test' (744 / 265 / 277 images respectively) 7via the 'file_names_complete' mapping files shipped in the archive. 8NOTE: the associated publication reports 1,009 "semantic segmentations", a smaller curated subset; 9this loader exposes the full 1,286 labeled images shipped in the 'thickness_segmentation_data' archive. 10 11The data is located at https://doi.org/10.5281/zenodo.3626020, released under a CC-BY-4.0 license. 12 13This dataset is from the publication https://doi.org/10.1038/s42256-020-00247-1. 14Please cite it if you use this dataset for your research. 15""" 16 17import os 18from glob import glob 19from typing import Union, Tuple, List, Literal 20 21import numpy as np 22import pandas as pd 23from PIL import Image 24 25from torch.utils.data import Dataset, DataLoader 26 27import torch_em 28 29from .. import util 30 31 32URL = "https://zenodo.org/records/3626020/files/thickness_segmentation_data.tar.gz" 33CHECKSUM = "f0c3c0ce9f140e470f6579bc6ec634d719d5947726f8b38f48f04f88026cdaa8" 34 35SPLITS = ["train", "validation", "test"] 36 37 38def _binarize_labels(data_dir): 39 label_paths = glob(os.path.join(data_dir, "data", "all_labels", "*.png")) 40 for label_path in label_paths: 41 label = np.array(Image.open(label_path).convert("L")) 42 Image.fromarray((label > 0).astype("uint8")).save(label_path) 43 44 45def get_deeprt_data(path: Union[os.PathLike, str], download: bool = False) -> str: 46 """Download the DeepRT dataset. 47 48 Args: 49 path: Filepath to a folder where the data is downloaded for further processing. 50 download: Whether to download the data if it is not present. 51 52 Returns: 53 Filepath where the data is downloaded. 54 """ 55 data_dir = os.path.join(path, "DeepRT") 56 if os.path.exists(data_dir): 57 return data_dir 58 59 os.makedirs(path, exist_ok=True) 60 61 tar_path = os.path.join(path, "thickness_segmentation_data.tar.gz") 62 util.download_source(path=tar_path, url=URL, download=download, checksum=CHECKSUM) 63 util.unzip_tarfile(tar_path=tar_path, dst=data_dir, remove=False) 64 65 _binarize_labels(data_dir) 66 67 return data_dir 68 69 70def get_deeprt_paths( 71 path: Union[os.PathLike, str], split: Literal["train", "validation", "test"] = "train", download: bool = False, 72) -> Tuple[List[str], List[str]]: 73 """Get paths to the DeepRT data. 74 75 Args: 76 path: Filepath to a folder where the data is downloaded for further processing. 77 split: The choice of data split. Either 'train', 'validation' or 'test'. 78 download: Whether to download the data if it is not present. 79 80 Returns: 81 List of filepaths for the image data. 82 List of filepaths for the label data. 83 """ 84 if split not in SPLITS: 85 raise ValueError(f"'{split}' is not a valid split. Choose one of {SPLITS}.") 86 87 data_dir = get_deeprt_data(path, download) 88 89 mapping = pd.read_csv(os.path.join(data_dir, "data", "file_names_complete", f"{split}_new_old_mapping.csv")) 90 ids = mapping["new_id"].tolist() 91 92 raw_paths = [os.path.join(data_dir, "data", "all_images", f"{i}.png") for i in ids] 93 label_paths = [os.path.join(data_dir, "data", "all_labels", f"{i}.png") for i in ids] 94 95 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 96 assert all(os.path.exists(p) for p in raw_paths) and all(os.path.exists(p) for p in label_paths) 97 98 return raw_paths, label_paths 99 100 101def get_deeprt_dataset( 102 path: Union[os.PathLike, str], 103 patch_shape: Tuple[int, int], 104 split: Literal["train", "validation", "test"] = "train", 105 resize_inputs: bool = False, 106 download: bool = False, 107 **kwargs 108) -> Dataset: 109 """Get the DeepRT dataset for retinal tissue segmentation in OCT images. 110 111 Args: 112 path: Filepath to a folder where the data is downloaded for further processing. 113 patch_shape: The patch shape to use for training. 114 split: The choice of data split. Either 'train', 'validation' or 'test'. 115 resize_inputs: Whether to resize the inputs to the patch shape. 116 download: Whether to download the data if it is not present. 117 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 118 119 Returns: 120 The segmentation dataset. 121 """ 122 raw_paths, label_paths = get_deeprt_paths(path, split, download) 123 124 if resize_inputs: 125 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True} 126 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 127 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 128 ) 129 130 return torch_em.default_segmentation_dataset( 131 raw_paths=raw_paths, 132 raw_key=None, 133 label_paths=label_paths, 134 label_key=None, 135 is_seg_dataset=False, 136 patch_shape=patch_shape, 137 ndim=2, 138 **kwargs 139 ) 140 141 142def get_deeprt_loader( 143 path: Union[os.PathLike, str], 144 batch_size: int, 145 patch_shape: Tuple[int, int], 146 split: Literal["train", "validation", "test"] = "train", 147 resize_inputs: bool = False, 148 download: bool = False, 149 **kwargs 150) -> DataLoader: 151 """Get the DeepRT dataloader for retinal tissue segmentation in OCT images. 152 153 Args: 154 path: Filepath to a folder where the data is downloaded for further processing. 155 batch_size: The batch size for training. 156 patch_shape: The patch shape to use for training. 157 split: The choice of data split. Either 'train', 'validation' or 'test'. 158 resize_inputs: Whether to resize the inputs to the patch shape. 159 download: Whether to download the data if it is not present. 160 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 161 162 Returns: 163 The DataLoader. 164 """ 165 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 166 dataset = get_deeprt_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 167 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
46def get_deeprt_data(path: Union[os.PathLike, str], download: bool = False) -> str: 47 """Download the DeepRT dataset. 48 49 Args: 50 path: Filepath to a folder where the data is downloaded for further processing. 51 download: Whether to download the data if it is not present. 52 53 Returns: 54 Filepath where the data is downloaded. 55 """ 56 data_dir = os.path.join(path, "DeepRT") 57 if os.path.exists(data_dir): 58 return data_dir 59 60 os.makedirs(path, exist_ok=True) 61 62 tar_path = os.path.join(path, "thickness_segmentation_data.tar.gz") 63 util.download_source(path=tar_path, url=URL, download=download, checksum=CHECKSUM) 64 util.unzip_tarfile(tar_path=tar_path, dst=data_dir, remove=False) 65 66 _binarize_labels(data_dir) 67 68 return data_dir
Download the DeepRT 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.
71def get_deeprt_paths( 72 path: Union[os.PathLike, str], split: Literal["train", "validation", "test"] = "train", download: bool = False, 73) -> Tuple[List[str], List[str]]: 74 """Get paths to the DeepRT data. 75 76 Args: 77 path: Filepath to a folder where the data is downloaded for further processing. 78 split: The choice of data split. Either 'train', 'validation' or 'test'. 79 download: Whether to download the data if it is not present. 80 81 Returns: 82 List of filepaths for the image data. 83 List of filepaths for the label data. 84 """ 85 if split not in SPLITS: 86 raise ValueError(f"'{split}' is not a valid split. Choose one of {SPLITS}.") 87 88 data_dir = get_deeprt_data(path, download) 89 90 mapping = pd.read_csv(os.path.join(data_dir, "data", "file_names_complete", f"{split}_new_old_mapping.csv")) 91 ids = mapping["new_id"].tolist() 92 93 raw_paths = [os.path.join(data_dir, "data", "all_images", f"{i}.png") for i in ids] 94 label_paths = [os.path.join(data_dir, "data", "all_labels", f"{i}.png") for i in ids] 95 96 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 97 assert all(os.path.exists(p) for p in raw_paths) and all(os.path.exists(p) for p in label_paths) 98 99 return raw_paths, label_paths
Get paths to the DeepRT data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The choice of data split. Either 'train', 'validation' or 'test'.
- 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.
102def get_deeprt_dataset( 103 path: Union[os.PathLike, str], 104 patch_shape: Tuple[int, int], 105 split: Literal["train", "validation", "test"] = "train", 106 resize_inputs: bool = False, 107 download: bool = False, 108 **kwargs 109) -> Dataset: 110 """Get the DeepRT dataset for retinal tissue segmentation in OCT images. 111 112 Args: 113 path: Filepath to a folder where the data is downloaded for further processing. 114 patch_shape: The patch shape to use for training. 115 split: The choice of data split. Either 'train', 'validation' or 'test'. 116 resize_inputs: Whether to resize the inputs to the patch shape. 117 download: Whether to download the data if it is not present. 118 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 119 120 Returns: 121 The segmentation dataset. 122 """ 123 raw_paths, label_paths = get_deeprt_paths(path, split, download) 124 125 if resize_inputs: 126 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True} 127 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 128 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 129 ) 130 131 return torch_em.default_segmentation_dataset( 132 raw_paths=raw_paths, 133 raw_key=None, 134 label_paths=label_paths, 135 label_key=None, 136 is_seg_dataset=False, 137 patch_shape=patch_shape, 138 ndim=2, 139 **kwargs 140 )
Get the DeepRT dataset for retinal tissue segmentation in OCT 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', 'validation' or 'test'.
- resize_inputs: Whether to resize the inputs to the 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.
143def get_deeprt_loader( 144 path: Union[os.PathLike, str], 145 batch_size: int, 146 patch_shape: Tuple[int, int], 147 split: Literal["train", "validation", "test"] = "train", 148 resize_inputs: bool = False, 149 download: bool = False, 150 **kwargs 151) -> DataLoader: 152 """Get the DeepRT dataloader for retinal tissue segmentation in OCT images. 153 154 Args: 155 path: Filepath to a folder where the data is downloaded for further processing. 156 batch_size: The batch size for training. 157 patch_shape: The patch shape to use for training. 158 split: The choice of data split. Either 'train', 'validation' or 'test'. 159 resize_inputs: Whether to resize the inputs to the patch shape. 160 download: Whether to download the data if it is not present. 161 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 162 163 Returns: 164 The DataLoader. 165 """ 166 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 167 dataset = get_deeprt_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 168 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the DeepRT dataloader for retinal tissue segmentation in OCT 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', 'validation' or 'test'.
- resize_inputs: Whether to resize the inputs to the 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.