torch_em.data.datasets.medical.semi_teethseg
The Semi-TeethSeg (STS-Tooth) dataset contains annotations for tooth segmentation in panoramic dental X-rays (PXI) and dental CBCT scans.
This is a multi-modal dataset for semi-supervised learning: alongside the labeled images, it also ships large unlabeled sets of panoramic X-rays and CBCT scans (not covered by this module, as they have no ground truth). This module covers both labeled subsets:
- The labeled panoramic X-ray subset (STS-2D-Tooth,
modality="2d"), with binary tooth segmentation masks for 900 adult and pediatric panoramic radiographs. - The labeled CBCT subset (STS-3D-Tooth,
modality="3d"), with 22 region-of-interest (ROI) volumes with per-tooth instance masks (subset="roi") and 10 whole field-of-view volumes with binary tooth masks (subset="integrity").
The dataset was curated for the MICCAI 2023 and 2024 Semi-supervised Teeth Segmentation (STS) challenges (https://sts-challenge.github.io/miccai2024/index.html) and is hosted on Zenodo at https://doi.org/10.5281/zenodo.10597292, distributed under the CC BY 4.0 license.
The dataset is from the publication https://doi.org/10.1038/s41597-024-04306-9. Please cite it if you use this dataset for your research.
1"""The Semi-TeethSeg (STS-Tooth) dataset contains annotations for tooth segmentation in panoramic 2dental X-rays (PXI) and dental CBCT scans. 3 4This is a multi-modal dataset for semi-supervised learning: alongside the labeled images, it also 5ships large unlabeled sets of panoramic X-rays and CBCT scans (not covered by this module, as they 6have no ground truth). This module covers both labeled subsets: 7- The labeled panoramic X-ray subset (STS-2D-Tooth, `modality="2d"`), with binary tooth segmentation 8 masks for 900 adult and pediatric panoramic radiographs. 9- The labeled CBCT subset (STS-3D-Tooth, `modality="3d"`), with 22 region-of-interest (ROI) volumes 10 with per-tooth instance masks (`subset="roi"`) and 10 whole field-of-view volumes with binary tooth 11 masks (`subset="integrity"`). 12 13The dataset was curated for the MICCAI 2023 and 2024 Semi-supervised Teeth Segmentation (STS) 14challenges (https://sts-challenge.github.io/miccai2024/index.html) and is hosted on Zenodo at 15https://doi.org/10.5281/zenodo.10597292, distributed under the CC BY 4.0 license. 16 17The dataset is from the publication https://doi.org/10.1038/s41597-024-04306-9. 18Please cite it if you use this dataset for your research. 19""" 20 21import os 22from glob import glob 23from tqdm import tqdm 24from pathlib import Path 25from natsort import natsorted 26from typing import Union, Literal, Tuple, List 27 28import numpy as np 29import imageio.v3 as imageio 30 31from torch.utils.data import Dataset, DataLoader 32 33import torch_em 34 35from .. import util 36 37 38BASE_URL = "https://zenodo.org/records/10597292/files" 39 40# The dataset is distributed as a 15-part split zip archive (SD-Tooth.zip.001 - .015, ~32GB total). 41# The parts are joined and unzipped once to extract the full 'SD-Tooth' folder. Only part 001 (which 42# already contains the full 'STS-2D-Tooth' subtree used by this module) has been downloaded and 43# verified so far, so only its checksum is set here; the others are left unverified. 44CHECKSUMS = { 45 "001": "916e9c9e72790c4bdf0884d012455c43821b2737bb27915e812dc20c2ff21cf4", 46 "002": None, 47 "003": None, 48 "004": None, 49 "005": None, 50 "006": None, 51 "007": None, 52 "008": None, 53 "009": None, 54 "010": None, 55 "011": None, 56 "012": None, 57 "013": None, 58 "014": None, 59 "015": None, 60} 61 62 63def get_semi_teethseg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 64 """Download the Semi-TeethSeg dataset. 65 66 Args: 67 path: Filepath to a folder where the data is downloaded for further processing. 68 download: Whether to download the data if it is not present. 69 70 Returns: 71 Filepath where the data is downloaded. 72 """ 73 data_dir = os.path.join(path, "SD-Tooth") 74 if os.path.exists(data_dir): 75 return data_dir 76 77 os.makedirs(path, exist_ok=True) 78 79 part_paths = [] 80 for part_id, checksum in CHECKSUMS.items(): 81 part_path = os.path.join(path, f"SD-Tooth.zip.{part_id}") 82 url = f"{BASE_URL}/SD-Tooth.zip.{part_id}?download=1" 83 util.download_source(path=part_path, url=url, download=download, checksum=checksum) 84 part_paths.append(part_path) 85 86 joined_zip_path = os.path.join(path, "SD-Tooth.zip") 87 if not os.path.exists(joined_zip_path): 88 with open(joined_zip_path, "wb") as dst: 89 for part_path in part_paths: 90 with open(part_path, "rb") as src: 91 while True: 92 chunk = src.read(1024 * 1024 * 64) 93 if not chunk: 94 break 95 dst.write(chunk) 96 97 util.unzip(zip_path=joined_zip_path, dst=path, remove=False) 98 99 return data_dir 100 101 102def get_semi_teethseg_paths( 103 path: Union[os.PathLike, str], 104 split: Literal["adult", "child"] = "adult", 105 modality: Literal["2d", "3d"] = "2d", 106 subset: Literal["roi", "integrity"] = "roi", 107 download: bool = False, 108) -> Tuple[List[str], List[str]]: 109 """Get paths to the Semi-TeethSeg data. 110 111 Args: 112 path: Filepath to a folder where the data is downloaded for further processing. 113 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 114 'child' (C-PXI). Only relevant for `modality="2d"`. 115 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 116 '3d' (CBCT volumes, STS-3D-Tooth). 117 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 118 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 119 tooth masks). Only relevant for `modality="3d"`. 120 download: Whether to download the data if it is not present. 121 122 Returns: 123 List of filepaths for the image data. 124 List of filepaths for the label data. 125 """ 126 data_dir = get_semi_teethseg_data(path, download) 127 128 if modality == "3d": 129 subset_dir = "ROI" if subset == "roi" else "Integrity" 130 base_dir = os.path.join(data_dir, "STS-3D-Tooth", subset_dir, "Labeled") 131 132 image_paths = natsorted(glob(os.path.join(base_dir, "Image", "*.nii.gz"))) 133 gt_paths = natsorted(glob(os.path.join(base_dir, "Mask", "*.nii.gz"))) 134 135 assert len(image_paths) == len(gt_paths) and len(image_paths) > 0 136 137 return image_paths, gt_paths 138 139 modality_dir = "A-PXI" if split == "adult" else "C-PXI" 140 base_dir = os.path.join(data_dir, "STS-2D-Tooth", modality_dir, "Labeled") 141 142 image_paths = natsorted(glob(os.path.join(base_dir, "Image", "*.png"))) 143 raw_gt_paths = natsorted(glob(os.path.join(base_dir, "Mask", "*.png"))) 144 145 assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0 146 147 neu_gt_dir = os.path.join(data_dir, "preprocessed", modality_dir) 148 os.makedirs(neu_gt_dir, exist_ok=True) 149 150 gt_paths = [] 151 for raw_gt_path in tqdm(raw_gt_paths, desc="Preprocessing labels"): 152 gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif") 153 gt_paths.append(gt_path) 154 if os.path.exists(gt_path): 155 continue 156 157 # the original masks are stored as single-bit (boolean) images, i.e. non-zero pixels 158 # correspond to teeth. we binarize them into a uint8 (0, 1) label map. 159 raw_gt = imageio.imread(raw_gt_path) 160 binary_gt = (raw_gt > 0).astype(np.uint8) 161 imageio.imwrite(gt_path, binary_gt) 162 163 return image_paths, gt_paths 164 165 166def get_semi_teethseg_dataset( 167 path: Union[os.PathLike, str], 168 patch_shape: Tuple[int, ...], 169 split: Literal["adult", "child"] = "adult", 170 modality: Literal["2d", "3d"] = "2d", 171 subset: Literal["roi", "integrity"] = "roi", 172 resize_inputs: bool = False, 173 download: bool = False, 174 **kwargs 175) -> Dataset: 176 """Get the Semi-TeethSeg dataset for tooth segmentation in panoramic dental radiographs or CBCT volumes. 177 178 Args: 179 path: Filepath to a folder where the data is downloaded for further processing. 180 patch_shape: The patch shape to use for training. 181 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 182 'child' (C-PXI). Only relevant for `modality="2d"`. 183 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 184 '3d' (CBCT volumes, STS-3D-Tooth). 185 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 186 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 187 tooth masks). Only relevant for `modality="3d"`. 188 resize_inputs: Whether to resize the inputs to the patch shape. 189 download: Whether to download the data if it is not present. 190 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 191 192 Returns: 193 The segmentation dataset. 194 """ 195 image_paths, gt_paths = get_semi_teethseg_paths(path, split, modality, subset, download) 196 197 if resize_inputs: 198 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 199 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 200 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 201 ) 202 203 if modality == "3d": 204 return torch_em.default_segmentation_dataset( 205 raw_paths=image_paths, 206 raw_key="data", 207 label_paths=gt_paths, 208 label_key="data", 209 is_seg_dataset=True, 210 patch_shape=patch_shape, 211 **kwargs 212 ) 213 214 return torch_em.default_segmentation_dataset( 215 raw_paths=image_paths, 216 raw_key=None, 217 label_paths=gt_paths, 218 label_key=None, 219 is_seg_dataset=False, 220 patch_shape=patch_shape, 221 **kwargs 222 ) 223 224 225def get_semi_teethseg_loader( 226 path: Union[os.PathLike, str], 227 batch_size: int, 228 patch_shape: Tuple[int, ...], 229 split: Literal["adult", "child"] = "adult", 230 modality: Literal["2d", "3d"] = "2d", 231 subset: Literal["roi", "integrity"] = "roi", 232 resize_inputs: bool = False, 233 download: bool = False, 234 **kwargs 235) -> DataLoader: 236 """Get the Semi-TeethSeg dataloader for tooth segmentation in panoramic dental radiographs or CBCT volumes. 237 238 Args: 239 path: Filepath to a folder where the data is downloaded for further processing. 240 batch_size: The batch size for training. 241 patch_shape: The patch shape to use for training. 242 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 243 'child' (C-PXI). Only relevant for `modality="2d"`. 244 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 245 '3d' (CBCT volumes, STS-3D-Tooth). 246 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 247 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 248 tooth masks). Only relevant for `modality="3d"`. 249 resize_inputs: Whether to resize the inputs to the patch shape. 250 download: Whether to download the data if it is not present. 251 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 252 253 Returns: 254 The DataLoader. 255 """ 256 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 257 dataset = get_semi_teethseg_dataset( 258 path, patch_shape, split, modality, subset, resize_inputs, download, **ds_kwargs 259 ) 260 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
64def get_semi_teethseg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 65 """Download the Semi-TeethSeg dataset. 66 67 Args: 68 path: Filepath to a folder where the data is downloaded for further processing. 69 download: Whether to download the data if it is not present. 70 71 Returns: 72 Filepath where the data is downloaded. 73 """ 74 data_dir = os.path.join(path, "SD-Tooth") 75 if os.path.exists(data_dir): 76 return data_dir 77 78 os.makedirs(path, exist_ok=True) 79 80 part_paths = [] 81 for part_id, checksum in CHECKSUMS.items(): 82 part_path = os.path.join(path, f"SD-Tooth.zip.{part_id}") 83 url = f"{BASE_URL}/SD-Tooth.zip.{part_id}?download=1" 84 util.download_source(path=part_path, url=url, download=download, checksum=checksum) 85 part_paths.append(part_path) 86 87 joined_zip_path = os.path.join(path, "SD-Tooth.zip") 88 if not os.path.exists(joined_zip_path): 89 with open(joined_zip_path, "wb") as dst: 90 for part_path in part_paths: 91 with open(part_path, "rb") as src: 92 while True: 93 chunk = src.read(1024 * 1024 * 64) 94 if not chunk: 95 break 96 dst.write(chunk) 97 98 util.unzip(zip_path=joined_zip_path, dst=path, remove=False) 99 100 return data_dir
Download the Semi-TeethSeg 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.
103def get_semi_teethseg_paths( 104 path: Union[os.PathLike, str], 105 split: Literal["adult", "child"] = "adult", 106 modality: Literal["2d", "3d"] = "2d", 107 subset: Literal["roi", "integrity"] = "roi", 108 download: bool = False, 109) -> Tuple[List[str], List[str]]: 110 """Get paths to the Semi-TeethSeg data. 111 112 Args: 113 path: Filepath to a folder where the data is downloaded for further processing. 114 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 115 'child' (C-PXI). Only relevant for `modality="2d"`. 116 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 117 '3d' (CBCT volumes, STS-3D-Tooth). 118 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 119 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 120 tooth masks). Only relevant for `modality="3d"`. 121 download: Whether to download the data if it is not present. 122 123 Returns: 124 List of filepaths for the image data. 125 List of filepaths for the label data. 126 """ 127 data_dir = get_semi_teethseg_data(path, download) 128 129 if modality == "3d": 130 subset_dir = "ROI" if subset == "roi" else "Integrity" 131 base_dir = os.path.join(data_dir, "STS-3D-Tooth", subset_dir, "Labeled") 132 133 image_paths = natsorted(glob(os.path.join(base_dir, "Image", "*.nii.gz"))) 134 gt_paths = natsorted(glob(os.path.join(base_dir, "Mask", "*.nii.gz"))) 135 136 assert len(image_paths) == len(gt_paths) and len(image_paths) > 0 137 138 return image_paths, gt_paths 139 140 modality_dir = "A-PXI" if split == "adult" else "C-PXI" 141 base_dir = os.path.join(data_dir, "STS-2D-Tooth", modality_dir, "Labeled") 142 143 image_paths = natsorted(glob(os.path.join(base_dir, "Image", "*.png"))) 144 raw_gt_paths = natsorted(glob(os.path.join(base_dir, "Mask", "*.png"))) 145 146 assert len(image_paths) == len(raw_gt_paths) and len(image_paths) > 0 147 148 neu_gt_dir = os.path.join(data_dir, "preprocessed", modality_dir) 149 os.makedirs(neu_gt_dir, exist_ok=True) 150 151 gt_paths = [] 152 for raw_gt_path in tqdm(raw_gt_paths, desc="Preprocessing labels"): 153 gt_path = os.path.join(neu_gt_dir, f"{Path(raw_gt_path).stem}.tif") 154 gt_paths.append(gt_path) 155 if os.path.exists(gt_path): 156 continue 157 158 # the original masks are stored as single-bit (boolean) images, i.e. non-zero pixels 159 # correspond to teeth. we binarize them into a uint8 (0, 1) label map. 160 raw_gt = imageio.imread(raw_gt_path) 161 binary_gt = (raw_gt > 0).astype(np.uint8) 162 imageio.imwrite(gt_path, binary_gt) 163 164 return image_paths, gt_paths
Get paths to the Semi-TeethSeg data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or
'child' (C-PXI). Only relevant for
modality="2d". - modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or '3d' (CBCT volumes, STS-3D-Tooth).
- subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with
per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary
tooth masks). Only relevant for
modality="3d". - 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.
167def get_semi_teethseg_dataset( 168 path: Union[os.PathLike, str], 169 patch_shape: Tuple[int, ...], 170 split: Literal["adult", "child"] = "adult", 171 modality: Literal["2d", "3d"] = "2d", 172 subset: Literal["roi", "integrity"] = "roi", 173 resize_inputs: bool = False, 174 download: bool = False, 175 **kwargs 176) -> Dataset: 177 """Get the Semi-TeethSeg dataset for tooth segmentation in panoramic dental radiographs or CBCT volumes. 178 179 Args: 180 path: Filepath to a folder where the data is downloaded for further processing. 181 patch_shape: The patch shape to use for training. 182 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 183 'child' (C-PXI). Only relevant for `modality="2d"`. 184 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 185 '3d' (CBCT volumes, STS-3D-Tooth). 186 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 187 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 188 tooth masks). Only relevant for `modality="3d"`. 189 resize_inputs: Whether to resize the inputs to the patch shape. 190 download: Whether to download the data if it is not present. 191 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 192 193 Returns: 194 The segmentation dataset. 195 """ 196 image_paths, gt_paths = get_semi_teethseg_paths(path, split, modality, subset, download) 197 198 if resize_inputs: 199 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 200 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 201 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 202 ) 203 204 if modality == "3d": 205 return torch_em.default_segmentation_dataset( 206 raw_paths=image_paths, 207 raw_key="data", 208 label_paths=gt_paths, 209 label_key="data", 210 is_seg_dataset=True, 211 patch_shape=patch_shape, 212 **kwargs 213 ) 214 215 return torch_em.default_segmentation_dataset( 216 raw_paths=image_paths, 217 raw_key=None, 218 label_paths=gt_paths, 219 label_key=None, 220 is_seg_dataset=False, 221 patch_shape=patch_shape, 222 **kwargs 223 )
Get the Semi-TeethSeg dataset for tooth segmentation in panoramic dental radiographs or CBCT volumes.
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 for the panoramic X-ray subset. Either 'adult' (A-PXI) or
'child' (C-PXI). Only relevant for
modality="2d". - modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or '3d' (CBCT volumes, STS-3D-Tooth).
- subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with
per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary
tooth masks). Only relevant for
modality="3d". - 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.
226def get_semi_teethseg_loader( 227 path: Union[os.PathLike, str], 228 batch_size: int, 229 patch_shape: Tuple[int, ...], 230 split: Literal["adult", "child"] = "adult", 231 modality: Literal["2d", "3d"] = "2d", 232 subset: Literal["roi", "integrity"] = "roi", 233 resize_inputs: bool = False, 234 download: bool = False, 235 **kwargs 236) -> DataLoader: 237 """Get the Semi-TeethSeg dataloader for tooth segmentation in panoramic dental radiographs or CBCT volumes. 238 239 Args: 240 path: Filepath to a folder where the data is downloaded for further processing. 241 batch_size: The batch size for training. 242 patch_shape: The patch shape to use for training. 243 split: The data split to use for the panoramic X-ray subset. Either 'adult' (A-PXI) or 244 'child' (C-PXI). Only relevant for `modality="2d"`. 245 modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or 246 '3d' (CBCT volumes, STS-3D-Tooth). 247 subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with 248 per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary 249 tooth masks). Only relevant for `modality="3d"`. 250 resize_inputs: Whether to resize the inputs to the patch shape. 251 download: Whether to download the data if it is not present. 252 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 253 254 Returns: 255 The DataLoader. 256 """ 257 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 258 dataset = get_semi_teethseg_dataset( 259 path, patch_shape, split, modality, subset, resize_inputs, download, **ds_kwargs 260 ) 261 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the Semi-TeethSeg dataloader for tooth segmentation in panoramic dental radiographs or CBCT volumes.
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 for the panoramic X-ray subset. Either 'adult' (A-PXI) or
'child' (C-PXI). Only relevant for
modality="2d". - modality: The choice of modality. Either '2d' (panoramic X-rays, STS-2D-Tooth) or '3d' (CBCT volumes, STS-3D-Tooth).
- subset: The choice of CBCT subset. Either 'roi' (22 region-of-interest volumes with
per-tooth instance masks) or 'integrity' (10 whole field-of-view volumes with binary
tooth masks). Only relevant for
modality="3d". - 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.