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)
URL = 'https://zenodo.org/records/3626020/files/thickness_segmentation_data.tar.gz'
CHECKSUM = 'f0c3c0ce9f140e470f6579bc6ec634d719d5947726f8b38f48f04f88026cdaa8'
SPLITS = ['train', 'validation', 'test']
def get_deeprt_data(path: Union[os.PathLike, str], download: bool = False) -> str:
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.

def get_deeprt_paths( path: Union[os.PathLike, str], split: Literal['train', 'validation', 'test'] = 'train', download: bool = False) -> Tuple[List[str], List[str]]:
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.

def get_deeprt_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'validation', 'test'] = 'train', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
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.

def get_deeprt_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'validation', 'test'] = 'train', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
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_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.