torch_em.data.datasets.medical.dprd

DPRD is a dataset for caries segmentation in children's dental panoramic radiographs.

This dataset is part of the "Children's Dental Panoramic Radiographs Dataset", which is hosted on figshare at https://doi.org/10.6084/m9.figshare.21621705.v1 (part of the collection https://doi.org/10.6084/m9.figshare.c.6317013.v1) and distributed under the CC0 license. This module only makes use of the "Children's dental caries segmentation dataset" subset, which is the part of the archive with pixel-level segmentation masks for dental caries. The raw masks are RGB images with black background and a fixed color marking the caries region; this module collapses them to a single-channel binary label map, where 0 is background and 1 marks caries.

The dataset is from the publication https://doi.org/10.1038/s41597-023-02237-5. Please cite it if you use this dataset for your research.

  1"""DPRD is a dataset for caries segmentation in children's dental panoramic radiographs.
  2
  3This dataset is part of the "Children's Dental Panoramic Radiographs Dataset", which is hosted on
  4figshare at https://doi.org/10.6084/m9.figshare.21621705.v1 (part of the collection
  5https://doi.org/10.6084/m9.figshare.c.6317013.v1) and distributed under the CC0 license. This module
  6only makes use of the "Children's dental caries segmentation dataset" subset, which is the part of
  7the archive with pixel-level segmentation masks for dental caries. The raw masks are RGB images
  8with black background and a fixed color marking the caries region; this module collapses them to
  9a single-channel binary label map, where 0 is background and 1 marks caries.
 10
 11The dataset is from the publication https://doi.org/10.1038/s41597-023-02237-5.
 12Please cite it if you use this dataset for your research.
 13"""
 14
 15import os
 16from glob import glob
 17from tqdm import tqdm
 18from pathlib import Path
 19from natsort import natsorted
 20from typing import Union, Tuple, List, Literal
 21
 22import numpy as np
 23import imageio.v3 as imageio
 24
 25from torch.utils.data import Dataset, DataLoader
 26
 27import torch_em
 28
 29from .. import util
 30
 31
 32URL = "https://ndownloader.figshare.com/files/38322366"
 33CHECKSUM = "2e41d862a0787828d659cfd13035e0fbaeb8687995dbfe29a84fdeac09b83945"
 34
 35ZIP_SUBDIR = "Children's dental caries segmentation dataset"
 36
 37
 38def get_dprd_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 39    """Download the DPRD 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, ZIP_SUBDIR)
 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, "Dental_dataset.zip")
 55    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
 56
 57    import zipfile
 58    with zipfile.ZipFile(zip_path) as f:
 59        members = [m for m in f.namelist() if m.startswith(f"{ZIP_SUBDIR}/")]
 60        f.extractall(path, members=members)
 61    os.remove(zip_path)
 62
 63    return data_dir
 64
 65
 66def get_dprd_paths(
 67    path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False
 68) -> Tuple[List[str], List[str]]:
 69    """Get paths to the DPRD data.
 70
 71    Args:
 72        path: Filepath to a folder where the data is downloaded for further processing.
 73        split: The data split to use. Either 'train' or 'test'.
 74        download: Whether to download the data if it is not present.
 75
 76    Returns:
 77        List of filepaths for the image data.
 78        List of filepaths for the label data.
 79    """
 80    if split not in ("train", "test"):
 81        raise ValueError(f"'{split}' is not a valid split. Please choose either 'train' or 'test'.")
 82
 83    data_dir = get_dprd_data(path, download)
 84
 85    split_dir = "Train" if split == "train" else "Test"
 86    image_paths = natsorted(glob(os.path.join(data_dir, split_dir, "images", "*.png")))
 87    raw_gt_paths = natsorted(glob(os.path.join(data_dir, split_dir, "mask", "*.png")))
 88
 89    assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0
 90
 91    neu_gt_dir = os.path.join(data_dir, "preprocessed", split)
 92    os.makedirs(neu_gt_dir, exist_ok=True)
 93
 94    gt_paths = []
 95    for raw_gt_path in tqdm(raw_gt_paths, desc="Preprocessing labels"):
 96        gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif")
 97        gt_paths.append(gt_path)
 98        if os.path.exists(gt_path):
 99            continue
100
101        # The raw masks are RGB images with black background and a fixed color (53, 119, 181)
102        # marking the caries region. We collapse this to a single-channel binary label map,
103        # where 0 is background and 1 marks caries.
104        raw_gt = imageio.imread(raw_gt_path)
105        binary_gt = (raw_gt.sum(axis=-1) > 0).astype(np.uint8)
106        imageio.imwrite(gt_path, binary_gt)
107
108    return image_paths, gt_paths
109
110
111def get_dprd_dataset(
112    path: Union[os.PathLike, str],
113    patch_shape: Tuple[int, int],
114    split: Literal["train", "test"],
115    resize_inputs: bool = False,
116    download: bool = False,
117    **kwargs
118) -> Dataset:
119    """Get the DPRD dataset for caries segmentation in panoramic dental radiographs.
120
121    Args:
122        path: Filepath to a folder where the data is downloaded for further processing.
123        patch_shape: The patch shape to use for training.
124        split: The data split to use. Either 'train' or 'test'.
125        resize_inputs: Whether to resize the inputs to the patch shape.
126        download: Whether to download the data if it is not present.
127        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
128
129    Returns:
130        The segmentation dataset.
131    """
132    image_paths, gt_paths = get_dprd_paths(path, split, download)
133
134    if resize_inputs:
135        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
136        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
137            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
138        )
139
140    return torch_em.default_segmentation_dataset(
141        raw_paths=image_paths,
142        raw_key=None,
143        label_paths=gt_paths,
144        label_key=None,
145        is_seg_dataset=False,
146        patch_shape=patch_shape,
147        **kwargs
148    )
149
150
151def get_dprd_loader(
152    path: Union[os.PathLike, str],
153    batch_size: int,
154    patch_shape: Tuple[int, int],
155    split: Literal["train", "test"],
156    resize_inputs: bool = False,
157    download: bool = False,
158    **kwargs
159) -> DataLoader:
160    """Get the DPRD dataloader for caries segmentation in panoramic dental radiographs.
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        split: The data split to use. Either 'train' or 'test'.
167        resize_inputs: Whether to resize the inputs to the 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_dprd_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
176    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://ndownloader.figshare.com/files/38322366'
CHECKSUM = '2e41d862a0787828d659cfd13035e0fbaeb8687995dbfe29a84fdeac09b83945'
ZIP_SUBDIR = "Children's dental caries segmentation dataset"
def get_dprd_data(path: Union[os.PathLike, str], download: bool = False) -> str:
39def get_dprd_data(path: Union[os.PathLike, str], download: bool = False) -> str:
40    """Download the DPRD dataset.
41
42    Args:
43        path: Filepath to a folder where the data is downloaded for further processing.
44        download: Whether to download the data if it is not present.
45
46    Returns:
47        Filepath where the data is downloaded.
48    """
49    data_dir = os.path.join(path, ZIP_SUBDIR)
50    if os.path.exists(data_dir):
51        return data_dir
52
53    os.makedirs(path, exist_ok=True)
54
55    zip_path = os.path.join(path, "Dental_dataset.zip")
56    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
57
58    import zipfile
59    with zipfile.ZipFile(zip_path) as f:
60        members = [m for m in f.namelist() if m.startswith(f"{ZIP_SUBDIR}/")]
61        f.extractall(path, members=members)
62    os.remove(zip_path)
63
64    return data_dir

Download the DPRD 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_dprd_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'], download: bool = False) -> Tuple[List[str], List[str]]:
 67def get_dprd_paths(
 68    path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False
 69) -> Tuple[List[str], List[str]]:
 70    """Get paths to the DPRD data.
 71
 72    Args:
 73        path: Filepath to a folder where the data is downloaded for further processing.
 74        split: The data split to use. Either 'train' or 'test'.
 75        download: Whether to download the data if it is not present.
 76
 77    Returns:
 78        List of filepaths for the image data.
 79        List of filepaths for the label data.
 80    """
 81    if split not in ("train", "test"):
 82        raise ValueError(f"'{split}' is not a valid split. Please choose either 'train' or 'test'.")
 83
 84    data_dir = get_dprd_data(path, download)
 85
 86    split_dir = "Train" if split == "train" else "Test"
 87    image_paths = natsorted(glob(os.path.join(data_dir, split_dir, "images", "*.png")))
 88    raw_gt_paths = natsorted(glob(os.path.join(data_dir, split_dir, "mask", "*.png")))
 89
 90    assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0
 91
 92    neu_gt_dir = os.path.join(data_dir, "preprocessed", split)
 93    os.makedirs(neu_gt_dir, exist_ok=True)
 94
 95    gt_paths = []
 96    for raw_gt_path in tqdm(raw_gt_paths, desc="Preprocessing labels"):
 97        gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif")
 98        gt_paths.append(gt_path)
 99        if os.path.exists(gt_path):
100            continue
101
102        # The raw masks are RGB images with black background and a fixed color (53, 119, 181)
103        # marking the caries region. We collapse this to a single-channel binary label map,
104        # where 0 is background and 1 marks caries.
105        raw_gt = imageio.imread(raw_gt_path)
106        binary_gt = (raw_gt.sum(axis=-1) > 0).astype(np.uint8)
107        imageio.imwrite(gt_path, binary_gt)
108
109    return image_paths, gt_paths

Get paths to the DPRD data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The data split to use. Either 'train' 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_dprd_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
112def get_dprd_dataset(
113    path: Union[os.PathLike, str],
114    patch_shape: Tuple[int, int],
115    split: Literal["train", "test"],
116    resize_inputs: bool = False,
117    download: bool = False,
118    **kwargs
119) -> Dataset:
120    """Get the DPRD dataset for caries segmentation in panoramic dental radiographs.
121
122    Args:
123        path: Filepath to a folder where the data is downloaded for further processing.
124        patch_shape: The patch shape to use for training.
125        split: The data split to use. Either 'train' or 'test'.
126        resize_inputs: Whether to resize the inputs to the 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, gt_paths = get_dprd_paths(path, split, 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=gt_paths,
145        label_key=None,
146        is_seg_dataset=False,
147        patch_shape=patch_shape,
148        **kwargs
149    )

Get the DPRD dataset for caries segmentation in panoramic dental radiographs.

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 data split to use. Either 'train' 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_dprd_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
152def get_dprd_loader(
153    path: Union[os.PathLike, str],
154    batch_size: int,
155    patch_shape: Tuple[int, int],
156    split: Literal["train", "test"],
157    resize_inputs: bool = False,
158    download: bool = False,
159    **kwargs
160) -> DataLoader:
161    """Get the DPRD dataloader for caries segmentation in panoramic dental radiographs.
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        split: The data split to use. Either 'train' or 'test'.
168        resize_inputs: Whether to resize the inputs to the patch shape.
169        download: Whether to download the data if it is not present.
170        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
171
172    Returns:
173        The DataLoader.
174    """
175    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
176    dataset = get_dprd_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
177    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the DPRD dataloader for caries segmentation in panoramic dental radiographs.

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 data split to use. Either 'train' 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.