torch_em.data.datasets.medical.aidk

AIDK is a dataset for the segmentation of the cornea, keratitis lesions and the iris in anterior-segment optical coherence tomography (AS-OCT) images.

The dataset contains 1,168 AS-OCT images from 64 keratitis patients: 400 'partial-frame' images (annotated for cornea and lesion) and 768 'full-frame' images (annotated for cornea, lesion and iris).

The dataset is located at https://doi.org/10.6084/m9.figshare.c.7036994.v1 and is distributed under the CC0 license. The dataset is from the publication https://doi.org/10.1038/s41597-024-03464-0. Please cite it if you use this dataset for your research.

  1"""AIDK is a dataset for the segmentation of the cornea, keratitis lesions and the iris in
  2anterior-segment optical coherence tomography (AS-OCT) images.
  3
  4The dataset contains 1,168 AS-OCT images from 64 keratitis patients: 400 'partial-frame' images
  5(annotated for cornea and lesion) and 768 'full-frame' images (annotated for cornea, lesion and iris).
  6
  7The dataset is located at https://doi.org/10.6084/m9.figshare.c.7036994.v1 and is distributed under
  8the CC0 license. The dataset is from the publication https://doi.org/10.1038/s41597-024-03464-0.
  9Please cite it if you use this dataset for your research.
 10"""
 11
 12import os
 13import json
 14from glob import glob
 15from tqdm import tqdm
 16from pathlib import Path
 17from natsort import natsorted
 18from typing import Union, Tuple, Literal, List
 19
 20import numpy as np
 21from skimage import draw
 22import imageio.v3 as imageio
 23
 24import torch_em
 25
 26from .. import util
 27
 28
 29URL = "https://ndownloader.figshare.com/files/46760137"
 30CHECKSUM = "70a0227324c288662fb8b02cc67a40facfc4615b085bf9631a965c1f8ed9c921"
 31
 32FRAMES = ["partial", "full"]
 33TASKS = ["cornea", "lesion", "iris"]
 34LABELS = {"cornea": "Cornea", "lesion": "Lesion", "iris": "Iris"}
 35
 36
 37def get_aidk_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 38    """Download the AIDK dataset.
 39
 40    Args:
 41        path: Filepath to a folder where the data is downloaded for further processing.
 42        download: Whether to download the data if it is not present.
 43
 44    Returns:
 45        Filepath where the data is downloaded.
 46    """
 47    data_dir = os.path.join(path, "AIDK_Dataset")
 48    if os.path.exists(data_dir):
 49        return data_dir
 50
 51    os.makedirs(path, exist_ok=True)
 52
 53    zip_path = os.path.join(path, "AIDK_Dataset.zip")
 54    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
 55    util.unzip(zip_path=zip_path, dst=path)
 56
 57    return data_dir
 58
 59
 60def _create_mask(annotation_path, image_shape, label):
 61    with open(annotation_path) as f:
 62        annotation = json.load(f)
 63
 64    mask = np.zeros(image_shape[:2], dtype=np.uint8)
 65    for shape in annotation["shapes"]:
 66        if shape["label"] != label or shape["shape_type"] != "polygon":
 67            continue
 68
 69        points = np.array(shape["points"])
 70        rr, cc = draw.polygon(points[:, 1], points[:, 0], shape=mask.shape)
 71        mask[rr, cc] = 1
 72
 73    return mask
 74
 75
 76def _preprocess_labels(data_dir, task, frame):
 77    label = LABELS[task]
 78    # Only the full-frame images carry iris annotations.
 79    frame_names = ["full"] if task == "iris" else ([frame] if frame != "all" else FRAMES)
 80
 81    image_paths, gt_paths = [], []
 82    for frame_name in frame_names:
 83        frame_dir = os.path.join(data_dir, f"{frame_name.capitalize()}-frame_Dataset")
 84        image_dir = os.path.join(frame_dir, "Original_AS-OCT_Images")
 85        annotation_dir = os.path.join(frame_dir, "Experts_Annotations")
 86        gt_dir = os.path.join(frame_dir, f"masks_{task}")
 87        os.makedirs(gt_dir, exist_ok=True)
 88
 89        annotation_paths = natsorted(
 90            p for p in glob(os.path.join(annotation_dir, "*.json")) if not os.path.basename(p).startswith("._")
 91        )
 92
 93        for annotation_path in tqdm(annotation_paths, desc=f"Converting '{frame_name}-frame' annotations to masks"):
 94            image_id = Path(annotation_path).stem
 95            image_path = os.path.join(image_dir, f"{image_id}.bmp")
 96            if not os.path.exists(image_path):
 97                continue
 98
 99            gt_path = os.path.join(gt_dir, f"{image_id}.tif")
100            if not os.path.exists(gt_path):
101                image_shape = imageio.imread(image_path).shape
102                mask = _create_mask(annotation_path, image_shape, label)
103                imageio.imwrite(gt_path, mask)
104
105            image_paths.append(image_path)
106            gt_paths.append(gt_path)
107
108    return image_paths, gt_paths
109
110
111def get_aidk_paths(
112    path: Union[os.PathLike, str],
113    task: Literal["cornea", "lesion", "iris"] = "cornea",
114    frame: Literal["partial", "full", "all"] = "all",
115    download: bool = False,
116) -> Tuple[List[str], List[str]]:
117    """Get paths to the AIDK data.
118
119    Args:
120        path: Filepath to a folder where the data is downloaded for further processing.
121        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
122        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
123            when `task` is 'iris', as only the full-frame images have iris annotations.
124        download: Whether to download the data if it is not present.
125
126    Returns:
127        List of filepaths for the image data.
128        List of filepaths for the label data.
129    """
130    if task not in TASKS:
131        raise ValueError(f"'{task}' is not a valid task. Please choose one of {TASKS}.")
132    if frame not in FRAMES + ["all"]:
133        raise ValueError(f"'{frame}' is not a valid frame choice. Please choose one of {FRAMES + ['all']}.")
134
135    data_dir = get_aidk_data(path=path, download=download)
136    image_paths, gt_paths = _preprocess_labels(data_dir, task, frame)
137
138    return image_paths, gt_paths
139
140
141def get_aidk_dataset(
142    path: Union[os.PathLike, str],
143    patch_shape: Tuple[int, int],
144    task: Literal["cornea", "lesion", "iris"] = "cornea",
145    frame: Literal["partial", "full", "all"] = "all",
146    resize_inputs: bool = False,
147    download: bool = False,
148    **kwargs
149):
150    """Get the AIDK dataset for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.
151
152    Args:
153        path: Filepath to a folder where the downloaded data will be saved.
154        patch_shape: The patch shape to use for training.
155        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
156        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
157            when `task` is 'iris'.
158        resize_inputs: Whether to resize the inputs to the expected 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`.
161
162    Returns:
163        The segmentation dataset.
164    """
165    image_paths, gt_paths = get_aidk_paths(path, task, frame, download)
166
167    if resize_inputs:
168        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
169        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
170            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
171        )
172
173    return torch_em.default_segmentation_dataset(
174        raw_paths=image_paths,
175        raw_key=None,
176        label_paths=gt_paths,
177        label_key=None,
178        patch_shape=patch_shape,
179        is_seg_dataset=False,
180        **kwargs
181    )
182
183
184def get_aidk_loader(
185    path: Union[os.PathLike, str],
186    batch_size: int,
187    patch_shape: Tuple[int, int],
188    task: Literal["cornea", "lesion", "iris"] = "cornea",
189    frame: Literal["partial", "full", "all"] = "all",
190    resize_inputs: bool = False,
191    download: bool = False,
192    **kwargs
193):
194    """Get the AIDK dataloader for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.
195
196    Args:
197        path: Filepath to a folder where the downloaded data will be saved.
198        batch_size: The batch size for training.
199        patch_shape: The patch shape to use for training.
200        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
201        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
202            when `task` is 'iris'.
203        resize_inputs: Whether to resize the inputs to the expected patch shape.
204        download: Whether to download the data if it is not present.
205        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
206
207    Returns:
208        The DataLoader.
209    """
210    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
211    dataset = get_aidk_dataset(path, patch_shape, task, frame, resize_inputs, download, **ds_kwargs)
212    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://ndownloader.figshare.com/files/46760137'
CHECKSUM = '70a0227324c288662fb8b02cc67a40facfc4615b085bf9631a965c1f8ed9c921'
FRAMES = ['partial', 'full']
TASKS = ['cornea', 'lesion', 'iris']
LABELS = {'cornea': 'Cornea', 'lesion': 'Lesion', 'iris': 'Iris'}
def get_aidk_data(path: Union[os.PathLike, str], download: bool = False) -> str:
38def get_aidk_data(path: Union[os.PathLike, str], download: bool = False) -> str:
39    """Download the AIDK dataset.
40
41    Args:
42        path: Filepath to a folder where the data is downloaded for further processing.
43        download: Whether to download the data if it is not present.
44
45    Returns:
46        Filepath where the data is downloaded.
47    """
48    data_dir = os.path.join(path, "AIDK_Dataset")
49    if os.path.exists(data_dir):
50        return data_dir
51
52    os.makedirs(path, exist_ok=True)
53
54    zip_path = os.path.join(path, "AIDK_Dataset.zip")
55    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
56    util.unzip(zip_path=zip_path, dst=path)
57
58    return data_dir

Download the AIDK 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_aidk_paths( path: Union[os.PathLike, str], task: Literal['cornea', 'lesion', 'iris'] = 'cornea', frame: Literal['partial', 'full', 'all'] = 'all', download: bool = False) -> Tuple[List[str], List[str]]:
112def get_aidk_paths(
113    path: Union[os.PathLike, str],
114    task: Literal["cornea", "lesion", "iris"] = "cornea",
115    frame: Literal["partial", "full", "all"] = "all",
116    download: bool = False,
117) -> Tuple[List[str], List[str]]:
118    """Get paths to the AIDK data.
119
120    Args:
121        path: Filepath to a folder where the data is downloaded for further processing.
122        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
123        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
124            when `task` is 'iris', as only the full-frame images have iris annotations.
125        download: Whether to download the data if it is not present.
126
127    Returns:
128        List of filepaths for the image data.
129        List of filepaths for the label data.
130    """
131    if task not in TASKS:
132        raise ValueError(f"'{task}' is not a valid task. Please choose one of {TASKS}.")
133    if frame not in FRAMES + ["all"]:
134        raise ValueError(f"'{frame}' is not a valid frame choice. Please choose one of {FRAMES + ['all']}.")
135
136    data_dir = get_aidk_data(path=path, download=download)
137    image_paths, gt_paths = _preprocess_labels(data_dir, task, frame)
138
139    return image_paths, gt_paths

Get paths to the AIDK data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
  • frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full') when task is 'iris', as only the full-frame images have iris annotations.
  • 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_aidk_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], task: Literal['cornea', 'lesion', 'iris'] = 'cornea', frame: Literal['partial', 'full', 'all'] = 'all', resize_inputs: bool = False, download: bool = False, **kwargs):
142def get_aidk_dataset(
143    path: Union[os.PathLike, str],
144    patch_shape: Tuple[int, int],
145    task: Literal["cornea", "lesion", "iris"] = "cornea",
146    frame: Literal["partial", "full", "all"] = "all",
147    resize_inputs: bool = False,
148    download: bool = False,
149    **kwargs
150):
151    """Get the AIDK dataset for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.
152
153    Args:
154        path: Filepath to a folder where the downloaded data will be saved.
155        patch_shape: The patch shape to use for training.
156        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
157        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
158            when `task` is 'iris'.
159        resize_inputs: Whether to resize the inputs to the expected 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`.
162
163    Returns:
164        The segmentation dataset.
165    """
166    image_paths, gt_paths = get_aidk_paths(path, task, frame, download)
167
168    if resize_inputs:
169        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
170        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
171            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
172        )
173
174    return torch_em.default_segmentation_dataset(
175        raw_paths=image_paths,
176        raw_key=None,
177        label_paths=gt_paths,
178        label_key=None,
179        patch_shape=patch_shape,
180        is_seg_dataset=False,
181        **kwargs
182    )

Get the AIDK dataset for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.

Arguments:
  • path: Filepath to a folder where the downloaded data will be saved.
  • patch_shape: The patch shape to use for training.
  • task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
  • frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full') when task is 'iris'.
  • 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_aidk_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], task: Literal['cornea', 'lesion', 'iris'] = 'cornea', frame: Literal['partial', 'full', 'all'] = 'all', resize_inputs: bool = False, download: bool = False, **kwargs):
185def get_aidk_loader(
186    path: Union[os.PathLike, str],
187    batch_size: int,
188    patch_shape: Tuple[int, int],
189    task: Literal["cornea", "lesion", "iris"] = "cornea",
190    frame: Literal["partial", "full", "all"] = "all",
191    resize_inputs: bool = False,
192    download: bool = False,
193    **kwargs
194):
195    """Get the AIDK dataloader for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.
196
197    Args:
198        path: Filepath to a folder where the downloaded data will be saved.
199        batch_size: The batch size for training.
200        patch_shape: The patch shape to use for training.
201        task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
202        frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full')
203            when `task` is 'iris'.
204        resize_inputs: Whether to resize the inputs to the expected patch shape.
205        download: Whether to download the data if it is not present.
206        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
207
208    Returns:
209        The DataLoader.
210    """
211    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
212    dataset = get_aidk_dataset(path, patch_shape, task, frame, resize_inputs, download, **ds_kwargs)
213    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the AIDK dataloader for segmentation of the cornea, keratitis lesions and the iris in AS-OCT images.

Arguments:
  • path: Filepath to a folder where the downloaded data will be saved.
  • batch_size: The batch size for training.
  • patch_shape: The patch shape to use for training.
  • task: The choice of segmentation task. Either 'cornea', 'lesion' or 'iris'.
  • frame: The choice of image subset. Either 'partial', 'full' or 'all'. Ignored (forced to 'full') when task is 'iris'.
  • 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.