torch_em.data.datasets.medical.drishti_gs

The Drishti-GS dataset contains annotations for optic disc and optic cup segmentation in Fundus images, for the task of glaucoma assessment.

The original data is hosted at https://cvit.iiit.ac.in/projects/mip/drishti-gs/mip-dataset2/Home.php, but downloading it (both the training and the test split) requires registering a team on that site. This dataloader uses a mirror of the data hosted on Kaggle instead: https://www.kaggle.com/datasets/lokeshsaipureddi/drishtigs-retina-dataset-for-onh-segmentation

The label masks are soft consensus maps (from four experts) stored as grayscale images, with values scaled to [0, 255]. We binarize them here using a majority vote threshold (i.e. more than half of the experts agree on the pixel belonging to the optic disc / cup).

The dataset is from the publication https://doi.org/10.1109/isbi.2014.6867807. Please cite it if you use this dataset for your research.

  1"""The Drishti-GS dataset contains annotations for optic disc and optic cup
  2segmentation in Fundus images, for the task of glaucoma assessment.
  3
  4The original data is hosted at https://cvit.iiit.ac.in/projects/mip/drishti-gs/mip-dataset2/Home.php,
  5but downloading it (both the training and the test split) requires registering a team on that site.
  6This dataloader uses a mirror of the data hosted on Kaggle instead:
  7https://www.kaggle.com/datasets/lokeshsaipureddi/drishtigs-retina-dataset-for-onh-segmentation
  8
  9The label masks are soft consensus maps (from four experts) stored as grayscale images, with values
 10scaled to `[0, 255]`. We binarize them here using a majority vote threshold (i.e. more than half of
 11the experts agree on the pixel belonging to the optic disc / cup).
 12
 13The dataset is from the publication https://doi.org/10.1109/isbi.2014.6867807.
 14Please cite it if you use this dataset for your research.
 15"""
 16
 17import os
 18from glob import glob
 19from pathlib import Path
 20from typing import Union, Tuple, Literal, List
 21
 22import imageio.v3 as imageio
 23
 24from torch.utils.data import Dataset, DataLoader
 25
 26import torch_em
 27
 28from .. import util
 29
 30
 31KAGGLE_DATASET_NAME = "lokeshsaipureddi/drishtigs-retina-dataset-for-onh-segmentation"
 32
 33
 34def get_drishti_gs_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 35    """Download the Drishti-GS dataset.
 36
 37    Args:
 38        path: Filepath to a folder where the data is downloaded for further processing.
 39        download: Whether to download the data if it is not present.
 40
 41    Returns:
 42        Filepath where the data is downloaded.
 43    """
 44    if os.path.exists(os.path.join(path, "Training")) and os.path.exists(os.path.join(path, "Test")):
 45        return path
 46
 47    os.makedirs(path, exist_ok=True)
 48
 49    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET_NAME, download=download)
 50    zip_path = os.path.join(path, "drishtigs-retina-dataset-for-onh-segmentation.zip")
 51    util.unzip(zip_path=zip_path, dst=path)
 52
 53    # The zip archive ships the training and test splits nested inside an extra, timestamped
 54    # directory level, e.g. 'Training-<timestamp>/Training'. Flatten this to 'Training'/'Test'.
 55    train_dir = glob(os.path.join(path, "Training*"))[0]
 56    os.rename(os.path.join(train_dir, "Training"), os.path.join(path, "Training"))
 57    os.rmdir(train_dir)
 58
 59    test_dir = glob(os.path.join(path, "Test*"))[0]
 60    os.rename(os.path.join(test_dir, "Test"), os.path.join(path, "Test"))
 61    os.rmdir(test_dir)
 62
 63    return path
 64
 65
 66def _binarize_mask(soft_map_path, gt_dir):
 67    dst_path = os.path.join(gt_dir, Path(soft_map_path).stem + ".tif")
 68    if os.path.exists(dst_path):
 69        return dst_path
 70
 71    os.makedirs(gt_dir, exist_ok=True)
 72    soft_map = imageio.imread(soft_map_path)
 73    mask = (soft_map > 127).astype("uint8")
 74    imageio.imwrite(dst_path, mask)
 75    return dst_path
 76
 77
 78def get_drishti_gs_paths(
 79    path: Union[os.PathLike, str],
 80    split: Literal["train", "test"],
 81    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
 82    download: bool = False,
 83) -> Tuple[List[str], List[str]]:
 84    """Get paths to the Drishti-GS data.
 85
 86    Args:
 87        path: Filepath to a folder where the data is downloaded for further processing.
 88        split: The choice of data split.
 89        task: The choice of labels for the specific task.
 90        download: Whether to download the data if it is not present.
 91
 92    Returns:
 93        List of filepaths for the image data.
 94        List of filepaths for the label data.
 95    """
 96    data_dir = get_drishti_gs_data(path=path, download=download)
 97
 98    assert split in ["train", "test"], f"'{split}' is not a valid split."
 99    assert task in ["optic_disc", "optic_cup"], f"'{task}' is not a valid task."
100
101    split_dir = "Training" if split == "train" else "Test"
102    gt_root = "GT" if split == "train" else "Test_GT"
103    softmap_name = "ODsegSoftmap" if task == "optic_disc" else "cupsegSoftmap"
104
105    image_paths = sorted(glob(os.path.join(data_dir, split_dir, "Images", "*", "*.png")))
106
107    gt_dir = os.path.join(data_dir, split_dir, f"gt_{task}")
108    gt_paths = []
109    for image_path in image_paths:
110        stem = Path(image_path).stem
111        soft_map_path = os.path.join(data_dir, split_dir, gt_root, stem, "SoftMap", f"{stem}_{softmap_name}.png")
112        gt_paths.append(_binarize_mask(soft_map_path, gt_dir))
113
114    assert len(image_paths) == len(gt_paths) and len(image_paths) > 0
115
116    return image_paths, gt_paths
117
118
119def get_drishti_gs_dataset(
120    path: Union[os.PathLike, str],
121    patch_shape: Tuple[int, int],
122    split: Literal["train", "test"],
123    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
124    resize_inputs: bool = False,
125    download: bool = False,
126    **kwargs
127) -> Dataset:
128    """Get the Drishti-GS dataset for segmentation of optic disc and optic cup in 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.
134        task: The choice of labels for the specific task.
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_drishti_gs_paths(path, split, task, 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    )
159
160
161def get_drishti_gs_loader(
162    path: Union[os.PathLike, str],
163    batch_size: int,
164    patch_shape: Tuple[int, int],
165    split: Literal["train", "test"],
166    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
167    resize_inputs: bool = False,
168    download: bool = False,
169    **kwargs
170) -> DataLoader:
171    """Get the Drishti-GS dataloader for segmentation of optic disc and optic cup in 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.
178        task: The choice of labels for the specific task.
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_drishti_gs_dataset(path, patch_shape, split, task, resize_inputs, download, **ds_kwargs)
188    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
KAGGLE_DATASET_NAME = 'lokeshsaipureddi/drishtigs-retina-dataset-for-onh-segmentation'
def get_drishti_gs_data(path: Union[os.PathLike, str], download: bool = False) -> str:
35def get_drishti_gs_data(path: Union[os.PathLike, str], download: bool = False) -> str:
36    """Download the Drishti-GS dataset.
37
38    Args:
39        path: Filepath to a folder where the data is downloaded for further processing.
40        download: Whether to download the data if it is not present.
41
42    Returns:
43        Filepath where the data is downloaded.
44    """
45    if os.path.exists(os.path.join(path, "Training")) and os.path.exists(os.path.join(path, "Test")):
46        return path
47
48    os.makedirs(path, exist_ok=True)
49
50    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET_NAME, download=download)
51    zip_path = os.path.join(path, "drishtigs-retina-dataset-for-onh-segmentation.zip")
52    util.unzip(zip_path=zip_path, dst=path)
53
54    # The zip archive ships the training and test splits nested inside an extra, timestamped
55    # directory level, e.g. 'Training-<timestamp>/Training'. Flatten this to 'Training'/'Test'.
56    train_dir = glob(os.path.join(path, "Training*"))[0]
57    os.rename(os.path.join(train_dir, "Training"), os.path.join(path, "Training"))
58    os.rmdir(train_dir)
59
60    test_dir = glob(os.path.join(path, "Test*"))[0]
61    os.rename(os.path.join(test_dir, "Test"), os.path.join(path, "Test"))
62    os.rmdir(test_dir)
63
64    return path

Download the Drishti-GS 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_drishti_gs_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'], task: Literal['optic_disc', 'optic_cup'] = 'optic_disc', download: bool = False) -> Tuple[List[str], List[str]]:
 79def get_drishti_gs_paths(
 80    path: Union[os.PathLike, str],
 81    split: Literal["train", "test"],
 82    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
 83    download: bool = False,
 84) -> Tuple[List[str], List[str]]:
 85    """Get paths to the Drishti-GS data.
 86
 87    Args:
 88        path: Filepath to a folder where the data is downloaded for further processing.
 89        split: The choice of data split.
 90        task: The choice of labels for the specific task.
 91        download: Whether to download the data if it is not present.
 92
 93    Returns:
 94        List of filepaths for the image data.
 95        List of filepaths for the label data.
 96    """
 97    data_dir = get_drishti_gs_data(path=path, download=download)
 98
 99    assert split in ["train", "test"], f"'{split}' is not a valid split."
100    assert task in ["optic_disc", "optic_cup"], f"'{task}' is not a valid task."
101
102    split_dir = "Training" if split == "train" else "Test"
103    gt_root = "GT" if split == "train" else "Test_GT"
104    softmap_name = "ODsegSoftmap" if task == "optic_disc" else "cupsegSoftmap"
105
106    image_paths = sorted(glob(os.path.join(data_dir, split_dir, "Images", "*", "*.png")))
107
108    gt_dir = os.path.join(data_dir, split_dir, f"gt_{task}")
109    gt_paths = []
110    for image_path in image_paths:
111        stem = Path(image_path).stem
112        soft_map_path = os.path.join(data_dir, split_dir, gt_root, stem, "SoftMap", f"{stem}_{softmap_name}.png")
113        gt_paths.append(_binarize_mask(soft_map_path, gt_dir))
114
115    assert len(image_paths) == len(gt_paths) and len(image_paths) > 0
116
117    return image_paths, gt_paths

Get paths to the Drishti-GS data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split.
  • task: The choice of labels for the specific task.
  • 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_drishti_gs_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'], task: Literal['optic_disc', 'optic_cup'] = 'optic_disc', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
120def get_drishti_gs_dataset(
121    path: Union[os.PathLike, str],
122    patch_shape: Tuple[int, int],
123    split: Literal["train", "test"],
124    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
125    resize_inputs: bool = False,
126    download: bool = False,
127    **kwargs
128) -> Dataset:
129    """Get the Drishti-GS dataset for segmentation of optic disc and optic cup in fundus images.
130
131    Args:
132        path: Filepath to a folder where the data is downloaded for further processing.
133        patch_shape: The patch shape to use for training.
134        split: The choice of data split.
135        task: The choice of labels for the specific task.
136        resize_inputs: Whether to resize the inputs to the expected patch shape.
137        download: Whether to download the data if it is not present.
138        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
139
140    Returns:
141        The segmentation dataset.
142    """
143    image_paths, gt_paths = get_drishti_gs_paths(path, split, task, download)
144
145    if resize_inputs:
146        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
147        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
148            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
149        )
150
151    return torch_em.default_segmentation_dataset(
152        raw_paths=image_paths,
153        raw_key=None,
154        label_paths=gt_paths,
155        label_key=None,
156        patch_shape=patch_shape,
157        is_seg_dataset=False,
158        **kwargs
159    )

Get the Drishti-GS dataset for segmentation of optic disc and optic cup 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.
  • task: The choice of labels for the specific task.
  • 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_drishti_gs_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'test'], task: Literal['optic_disc', 'optic_cup'] = 'optic_disc', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
162def get_drishti_gs_loader(
163    path: Union[os.PathLike, str],
164    batch_size: int,
165    patch_shape: Tuple[int, int],
166    split: Literal["train", "test"],
167    task: Literal["optic_disc", "optic_cup"] = "optic_disc",
168    resize_inputs: bool = False,
169    download: bool = False,
170    **kwargs
171) -> DataLoader:
172    """Get the Drishti-GS dataloader for segmentation of optic disc and optic cup in fundus images.
173
174    Args:
175        path: Filepath to a folder where the data is downloaded for further processing.
176        batch_size: The batch size for training.
177        patch_shape: The patch shape to use for training.
178        split: The choice of data split.
179        task: The choice of labels for the specific task.
180        resize_inputs: Whether to resize the inputs to the expected patch shape.
181        download: Whether to download the data if it is not present.
182        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
183
184    Returns:
185        The DataLoader.
186    """
187    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
188    dataset = get_drishti_gs_dataset(path, patch_shape, split, task, resize_inputs, download, **ds_kwargs)
189    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the Drishti-GS dataloader for segmentation of optic disc and optic cup 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.
  • task: The choice of labels for the specific task.
  • 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.