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)
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.
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.
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.