torch_em.data.datasets.medical.hrf_seg_plus

The HRF-Seg+ dataset contains manual, expert-reviewed multi-structure annotations for the fundus images of the HRF dataset (torch_em.data.datasets.medical.hrf): the optic disc, the optic cup, the retinal vessels and the peripapillary alpha and beta zones.

NOTE: HRF-Seg+ ships its own copy of the 45 HRF fundus images (resized to 500x500), rather than reusing the original high-resolution HRF images directly, so this dataset is implemented as a standalone module rather than as an extension of torch_em.data.datasets.medical.hrf.

The archive stores five annotated structures as separate folders: 'Folder_1_Optic_Disc', 'Folder_2_Optic_Cup', 'Folder_3_Vessels' and 'Folder_4_Alpha_Beta_Zones' (each structure as an RGBA image, with the structure's pixels opaque), plus a merged multi-class mask per image in 'Folder_5_Ground_Truth_Masks' (an RGB image, colored per the 'class_dict.csv' palette). We use the merged masks and map each pixel to the class with the closest reference color (to be robust to the lossy-compression-free but anti-aliased mask boundaries), yielding semantic labels: 0 (background), 1 (optic disc), 2 (retinal vessels), 3 (optic cup), 4 (beta zone) and 5 (alpha zone).

This dataset is located at https://doi.org/10.5281/zenodo.16744782. The dataset is licensed under CC-BY-4.0. Please cite the dataset if you use it in your research.

  1"""The HRF-Seg+ dataset contains manual, expert-reviewed multi-structure annotations for the fundus
  2images of the HRF dataset (`torch_em.data.datasets.medical.hrf`): the optic disc, the optic cup, the
  3retinal vessels and the peripapillary alpha and beta zones.
  4
  5NOTE: HRF-Seg+ ships its own copy of the 45 HRF fundus images (resized to 500x500), rather than
  6reusing the original high-resolution HRF images directly, so this dataset is implemented as a
  7standalone module rather than as an extension of `torch_em.data.datasets.medical.hrf`.
  8
  9The archive stores five annotated structures as separate folders: 'Folder_1_Optic_Disc',
 10'Folder_2_Optic_Cup', 'Folder_3_Vessels' and 'Folder_4_Alpha_Beta_Zones' (each structure as an RGBA
 11image, with the structure's pixels opaque), plus a merged multi-class mask per image in
 12'Folder_5_Ground_Truth_Masks' (an RGB image, colored per the 'class_dict.csv' palette). We use the
 13merged masks and map each pixel to the class with the closest reference color (to be robust to the
 14lossy-compression-free but anti-aliased mask boundaries), yielding semantic labels: 0 (background),
 151 (optic disc), 2 (retinal vessels), 3 (optic cup), 4 (beta zone) and 5 (alpha zone).
 16
 17This dataset is located at https://doi.org/10.5281/zenodo.16744782.
 18The dataset is licensed under CC-BY-4.0.
 19Please cite the dataset if you use it in your research.
 20"""
 21
 22import os
 23from glob import glob
 24from tqdm import tqdm
 25from pathlib import Path
 26from natsort import natsorted
 27from typing import Union, Tuple, List
 28
 29import numpy as np
 30import imageio.v3 as imageio
 31
 32from torch.utils.data import Dataset, DataLoader
 33
 34import torch_em
 35
 36from .. import util
 37
 38
 39URL = "https://zenodo.org/records/16744782/files/HRF-Seg%2B.zip"
 40CHECKSUM = "e4bf59bb820147aa28158b1ed2228cf59cc426605d023ab2edca651fcf3f3411"
 41
 42LABEL_COLORS = {
 43    0: (0, 0, 0),  # unlabeled
 44    1: (128, 64, 128),  # opticdisc
 45    2: (254, 148, 12),  # retinalvessels
 46    3: (130, 76, 0),  # opticcup
 47    4: (190, 250, 190),  # betazone
 48    5: (112, 150, 146),  # alphazone
 49}
 50
 51
 52def get_hrf_seg_plus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 53    """Download the HRF-Seg+ dataset.
 54
 55    Args:
 56        path: Filepath to a folder where the data is downloaded for further processing.
 57        download: Whether to download the data if it is not present.
 58
 59    Returns:
 60        Filepath where the data is downloaded.
 61    """
 62    data_dir = os.path.join(path, "HRF-Seg+")
 63    if os.path.exists(data_dir):
 64        return data_dir
 65
 66    os.makedirs(path, exist_ok=True)
 67
 68    zip_path = os.path.join(path, "HRF-Seg+.zip")
 69    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
 70    util.unzip(zip_path=zip_path, dst=path)
 71
 72    return data_dir
 73
 74
 75def get_hrf_seg_plus_paths(
 76    path: Union[os.PathLike, str], download: bool = False
 77) -> Tuple[List[str], List[str]]:
 78    """Get paths to the HRF-Seg+ data.
 79
 80    Args:
 81        path: Filepath to a folder where the data is downloaded for further processing.
 82        download: Whether to download the data if it is not present.
 83
 84    Returns:
 85        List of filepaths for the image data.
 86        List of filepaths for the label data.
 87    """
 88    data_dir = get_hrf_seg_plus_data(path, download)
 89
 90    image_paths = natsorted(glob(os.path.join(data_dir, "Folder_6_Original_Images", "*")))
 91    mask_paths = natsorted(glob(os.path.join(data_dir, "Folder_5_Ground_Truth_Masks", "*.png")))
 92    assert image_paths and len(image_paths) == len(mask_paths)
 93
 94    label_dir = os.path.join(data_dir, "preprocessed_labels")
 95    os.makedirs(label_dir, exist_ok=True)
 96
 97    reference_colors = np.array(list(LABEL_COLORS.values()))
 98
 99    label_paths = []
100    for mask_path in tqdm(mask_paths, desc="Preprocessing labels"):
101        label_path = os.path.join(label_dir, f"{Path(mask_path).stem}.tif")
102        label_paths.append(label_path)
103        if os.path.exists(label_path):
104            continue
105
106        mask = imageio.imread(mask_path)[..., :3].astype("float32")
107        distances = np.linalg.norm(mask[..., None, :] - reference_colors[None, None, :, :], axis=-1)
108        semantic_mask = np.argmin(distances, axis=-1).astype("uint8")
109        imageio.imwrite(label_path, semantic_mask, compression="zlib")
110
111    return image_paths, label_paths
112
113
114def get_hrf_seg_plus_dataset(
115    path: Union[os.PathLike, str],
116    patch_shape: Tuple[int, int],
117    resize_inputs: bool = False,
118    download: bool = False,
119    **kwargs
120) -> Dataset:
121    """Get the HRF-Seg+ dataset for optic disc, optic cup, vessel and peripapillary zone segmentation.
122
123    Args:
124        path: Filepath to a folder where the data is downloaded for further processing.
125        patch_shape: The patch shape to use for training.
126        resize_inputs: Whether to resize the inputs to the expected patch shape.
127        download: Whether to download the data if it is not present.
128        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
129
130    Returns:
131        The segmentation dataset.
132    """
133    image_paths, label_paths = get_hrf_seg_plus_paths(path, download)
134
135    if resize_inputs:
136        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
137        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
138            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
139        )
140
141    return torch_em.default_segmentation_dataset(
142        raw_paths=image_paths,
143        raw_key=None,
144        label_paths=label_paths,
145        label_key=None,
146        patch_shape=patch_shape,
147        is_seg_dataset=False,
148        **kwargs
149    )
150
151
152def get_hrf_seg_plus_loader(
153    path: Union[os.PathLike, str],
154    batch_size: int,
155    patch_shape: Tuple[int, int],
156    resize_inputs: bool = False,
157    download: bool = False,
158    **kwargs
159) -> DataLoader:
160    """Get the HRF-Seg+ dataloader for optic disc, optic cup, vessel and peripapillary zone segmentation.
161
162    Args:
163        path: Filepath to a folder where the data is downloaded for further processing.
164        batch_size: The batch size for training.
165        patch_shape: The patch shape to use for training.
166        resize_inputs: Whether to resize the inputs to the expected patch shape.
167        download: Whether to download the data if it is not present.
168        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
169
170    Returns:
171        The DataLoader.
172    """
173    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
174    dataset = get_hrf_seg_plus_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
175    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://zenodo.org/records/16744782/files/HRF-Seg%2B.zip'
CHECKSUM = 'e4bf59bb820147aa28158b1ed2228cf59cc426605d023ab2edca651fcf3f3411'
LABEL_COLORS = {0: (0, 0, 0), 1: (128, 64, 128), 2: (254, 148, 12), 3: (130, 76, 0), 4: (190, 250, 190), 5: (112, 150, 146)}
def get_hrf_seg_plus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
53def get_hrf_seg_plus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
54    """Download the HRF-Seg+ dataset.
55
56    Args:
57        path: Filepath to a folder where the data is downloaded for further processing.
58        download: Whether to download the data if it is not present.
59
60    Returns:
61        Filepath where the data is downloaded.
62    """
63    data_dir = os.path.join(path, "HRF-Seg+")
64    if os.path.exists(data_dir):
65        return data_dir
66
67    os.makedirs(path, exist_ok=True)
68
69    zip_path = os.path.join(path, "HRF-Seg+.zip")
70    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
71    util.unzip(zip_path=zip_path, dst=path)
72
73    return data_dir

Download the HRF-Seg+ 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_hrf_seg_plus_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
 76def get_hrf_seg_plus_paths(
 77    path: Union[os.PathLike, str], download: bool = False
 78) -> Tuple[List[str], List[str]]:
 79    """Get paths to the HRF-Seg+ data.
 80
 81    Args:
 82        path: Filepath to a folder where the data is downloaded for further processing.
 83        download: Whether to download the data if it is not present.
 84
 85    Returns:
 86        List of filepaths for the image data.
 87        List of filepaths for the label data.
 88    """
 89    data_dir = get_hrf_seg_plus_data(path, download)
 90
 91    image_paths = natsorted(glob(os.path.join(data_dir, "Folder_6_Original_Images", "*")))
 92    mask_paths = natsorted(glob(os.path.join(data_dir, "Folder_5_Ground_Truth_Masks", "*.png")))
 93    assert image_paths and len(image_paths) == len(mask_paths)
 94
 95    label_dir = os.path.join(data_dir, "preprocessed_labels")
 96    os.makedirs(label_dir, exist_ok=True)
 97
 98    reference_colors = np.array(list(LABEL_COLORS.values()))
 99
100    label_paths = []
101    for mask_path in tqdm(mask_paths, desc="Preprocessing labels"):
102        label_path = os.path.join(label_dir, f"{Path(mask_path).stem}.tif")
103        label_paths.append(label_path)
104        if os.path.exists(label_path):
105            continue
106
107        mask = imageio.imread(mask_path)[..., :3].astype("float32")
108        distances = np.linalg.norm(mask[..., None, :] - reference_colors[None, None, :, :], axis=-1)
109        semantic_mask = np.argmin(distances, axis=-1).astype("uint8")
110        imageio.imwrite(label_path, semantic_mask, compression="zlib")
111
112    return image_paths, label_paths

Get paths to the HRF-Seg+ data.

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:

List of filepaths for the image data. List of filepaths for the label data.

def get_hrf_seg_plus_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
115def get_hrf_seg_plus_dataset(
116    path: Union[os.PathLike, str],
117    patch_shape: Tuple[int, int],
118    resize_inputs: bool = False,
119    download: bool = False,
120    **kwargs
121) -> Dataset:
122    """Get the HRF-Seg+ dataset for optic disc, optic cup, vessel and peripapillary zone segmentation.
123
124    Args:
125        path: Filepath to a folder where the data is downloaded for further processing.
126        patch_shape: The patch shape to use for training.
127        resize_inputs: Whether to resize the inputs to the expected patch shape.
128        download: Whether to download the data if it is not present.
129        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
130
131    Returns:
132        The segmentation dataset.
133    """
134    image_paths, label_paths = get_hrf_seg_plus_paths(path, download)
135
136    if resize_inputs:
137        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
138        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
139            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
140        )
141
142    return torch_em.default_segmentation_dataset(
143        raw_paths=image_paths,
144        raw_key=None,
145        label_paths=label_paths,
146        label_key=None,
147        patch_shape=patch_shape,
148        is_seg_dataset=False,
149        **kwargs
150    )

Get the HRF-Seg+ dataset for optic disc, optic cup, vessel and peripapillary zone segmentation.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • 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_hrf_seg_plus_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
153def get_hrf_seg_plus_loader(
154    path: Union[os.PathLike, str],
155    batch_size: int,
156    patch_shape: Tuple[int, int],
157    resize_inputs: bool = False,
158    download: bool = False,
159    **kwargs
160) -> DataLoader:
161    """Get the HRF-Seg+ dataloader for optic disc, optic cup, vessel and peripapillary zone segmentation.
162
163    Args:
164        path: Filepath to a folder where the data is downloaded for further processing.
165        batch_size: The batch size for training.
166        patch_shape: The patch shape to use for training.
167        resize_inputs: Whether to resize the inputs to the expected patch shape.
168        download: Whether to download the data if it is not present.
169        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
170
171    Returns:
172        The DataLoader.
173    """
174    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
175    dataset = get_hrf_seg_plus_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
176    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the HRF-Seg+ dataloader for optic disc, optic cup, vessel and peripapillary zone segmentation.

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.
  • 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.