torch_em.data.datasets.medical.pmcanalseg

The PMCanalSeg dataset contains annotations for segmentation of the maxillary pterygopalatine canal and the mandibular canal in 3D CBCT images.

The dataset is located at https://doi.org/10.7910/DVN/RTIGTP, hosted on Harvard Dataverse under a CC0 1.0 license.

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

The dataset comprises 191 patients, each with an 'upper' scan (maxilla) annotated for the pterygopalatine canal and a 'lower' scan (mandible) annotated for the mandibular canal, plus an unannotated full 'skull' scan. This module only exposes the annotated 'upper' and 'lower' volumes.

  1"""The PMCanalSeg dataset contains annotations for segmentation of the maxillary pterygopalatine
  2canal and the mandibular canal in 3D CBCT images.
  3
  4The dataset is located at https://doi.org/10.7910/DVN/RTIGTP, hosted on Harvard Dataverse under
  5a CC0 1.0 license.
  6
  7The dataset is from the publication https://doi.org/10.1038/s41597-026-06620-w.
  8Please cite it if you use this dataset for your research.
  9
 10The dataset comprises 191 patients, each with an 'upper' scan (maxilla) annotated for the
 11pterygopalatine canal and a 'lower' scan (mandible) annotated for the mandibular canal, plus an
 12unannotated full 'skull' scan. This module only exposes the annotated 'upper' and 'lower' volumes.
 13"""
 14
 15import os
 16import hashlib
 17from glob import glob
 18from natsort import natsorted
 19from typing import Union, Tuple, Literal, List
 20
 21import requests
 22from tqdm import tqdm
 23
 24from torch.utils.data import Dataset, DataLoader
 25
 26import torch_em
 27
 28from .. import util
 29
 30
 31PERSISTENT_ID = "doi:10.7910/DVN/RTIGTP"
 32BASE_URL = "https://dataverse.harvard.edu"
 33
 34# The Dataverse API rejects requests with the default 'python-requests' user agent (403 Forbidden).
 35HEADERS = {"User-Agent": "Mozilla/5.0"}
 36
 37
 38def _get_manifest(path):
 39    import json
 40
 41    manifest_path = os.path.join(path, "manifest.json")
 42    if os.path.exists(manifest_path):
 43        with open(manifest_path) as f:
 44            return json.load(f)
 45
 46    url = f"{BASE_URL}/api/datasets/:persistentId/versions/:latest?persistentId={PERSISTENT_ID}"
 47    r = requests.get(url, headers=HEADERS)
 48    r.raise_for_status()
 49    files = r.json()["data"]["files"]
 50
 51    manifest = []
 52    for f in files:
 53        directory_label = f.get("directoryLabel", "")
 54        if not directory_label.startswith(("upper/", "lower/")):
 55            continue
 56
 57        data_file = f["dataFile"]
 58        manifest.append({
 59            "directory": directory_label,
 60            "filename": data_file["filename"],
 61            "id": data_file["id"],
 62            "md5": data_file.get("md5"),
 63        })
 64
 65    os.makedirs(path, exist_ok=True)
 66    with open(manifest_path, "w") as f:
 67        json.dump(manifest, f)
 68
 69    return manifest
 70
 71
 72def _download_file(url, path, md5=None):
 73    if os.path.exists(path):
 74        return
 75
 76    tmp_path = f"{path}.incomplete"
 77    with requests.get(url, stream=True, headers=HEADERS) as r:
 78        r.raise_for_status()
 79        file_size = int(r.headers.get("Content-Length", 0))
 80        with tqdm.wrapattr(r.raw, "read", total=file_size, desc=f"Download {url} to {path}") as r_raw:
 81            with open(tmp_path, "wb") as f:
 82                for chunk in iter(lambda: r_raw.read(1 << 20), b""):
 83                    f.write(chunk)
 84
 85    if md5 is not None:
 86        hasher = hashlib.md5()
 87        with open(tmp_path, "rb") as f:
 88            for chunk in iter(lambda: f.read(1 << 20), b""):
 89                hasher.update(chunk)
 90        if hasher.hexdigest() != md5:
 91            raise RuntimeError(f"The checksum of {url} does not match the expected checksum.")
 92
 93    os.replace(tmp_path, path)
 94
 95
 96def get_pmcanalseg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 97    """Download the PMCanalSeg dataset.
 98
 99    Args:
100        path: Filepath to a folder where the data is downloaded for further processing.
101        download: Whether to download the data if it is not present.
102
103    Returns:
104        Filepath where the data is downloaded.
105    """
106    os.makedirs(path, exist_ok=True)
107    manifest = _get_manifest(path)
108
109    missing = [entry for entry in manifest if not os.path.exists(os.path.join(path, entry["directory"], entry["filename"]))]  # noqa
110    if missing and not download:
111        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
112
113    for entry in tqdm(missing, desc="Downloading PMCanalSeg"):
114        out_dir = os.path.join(path, entry["directory"])
115        os.makedirs(out_dir, exist_ok=True)
116        out_path = os.path.join(out_dir, entry["filename"])
117        url = f"{BASE_URL}/api/access/datafile/{entry['id']}"
118        _download_file(url, out_path, entry["md5"])
119
120    return path
121
122
123def get_pmcanalseg_paths(
124    path: Union[os.PathLike, str],
125    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
126    download: bool = False,
127) -> Tuple[List[str], List[str]]:
128    """Get paths to the PMCanalSeg data.
129
130    Args:
131        path: Filepath to a folder where the data is downloaded for further processing.
132        label_choice: The choice of canal to segment. Either 'mandibular' (from the 'lower'
133            mandible CBCT scans) or 'pterygopalatine' (from the 'upper' maxillary CBCT scans).
134        download: Whether to download the data if it is not present.
135
136    Returns:
137        List of filepaths for the image data.
138        List of filepaths for the label data.
139    """
140    if label_choice not in ("mandibular", "pterygopalatine"):
141        raise ValueError(f"'{label_choice}' is not a valid label choice. Please choose 'mandibular' or 'pterygopalatine'.")  # noqa
142
143    data_dir = get_pmcanalseg_data(path, download)
144    subdir = "lower" if label_choice == "mandibular" else "upper"
145
146    image_paths = natsorted(glob(os.path.join(data_dir, subdir, "Patient_*", "image.nii.gz")))
147    gt_paths = [p.replace("image.nii.gz", "label.nii.gz") for p in image_paths]
148
149    image_paths = [p for p, g in zip(image_paths, gt_paths) if os.path.exists(g)]
150    gt_paths = [g for g in gt_paths if os.path.exists(g)]
151
152    return image_paths, gt_paths
153
154
155def get_pmcanalseg_dataset(
156    path: Union[os.PathLike, str],
157    patch_shape: Tuple[int, ...],
158    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
159    resize_inputs: bool = False,
160    download: bool = False,
161    **kwargs
162) -> Dataset:
163    """Get the PMCanalSeg dataset for segmentation of the mandibular or pterygopalatine canal in CBCT.
164
165    Args:
166        path: Filepath to a folder where the data is downloaded for further processing.
167        patch_shape: The patch shape to use for training.
168        label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
169        resize_inputs: Whether to resize the inputs to the patch shape.
170        download: Whether to download the data if it is not present.
171        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
172
173    Returns:
174        The segmentation dataset.
175    """
176    image_paths, gt_paths = get_pmcanalseg_paths(path, label_choice, download)
177
178    if resize_inputs:
179        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
180        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
181            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
182        )
183
184    return torch_em.default_segmentation_dataset(
185        raw_paths=image_paths,
186        raw_key="data",
187        label_paths=gt_paths,
188        label_key="data",
189        patch_shape=patch_shape,
190        is_seg_dataset=True,
191        **kwargs
192    )
193
194
195def get_pmcanalseg_loader(
196    path: Union[os.PathLike, str],
197    batch_size: int,
198    patch_shape: Tuple[int, ...],
199    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
200    resize_inputs: bool = False,
201    download: bool = False,
202    **kwargs
203) -> DataLoader:
204    """Get the PMCanalSeg dataloader for segmentation of the mandibular or pterygopalatine canal in CBCT.
205
206    Args:
207        path: Filepath to a folder where the data is downloaded for further processing.
208        batch_size: The batch size for training.
209        patch_shape: The patch shape to use for training.
210        label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
211        resize_inputs: Whether to resize the inputs to the patch shape.
212        download: Whether to download the data if it is not present.
213        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
214
215    Returns:
216        The DataLoader.
217    """
218    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
219    dataset = get_pmcanalseg_dataset(path, patch_shape, label_choice, resize_inputs, download, **ds_kwargs)
220    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
PERSISTENT_ID = 'doi:10.7910/DVN/RTIGTP'
BASE_URL = 'https://dataverse.harvard.edu'
HEADERS = {'User-Agent': 'Mozilla/5.0'}
def get_pmcanalseg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 97def get_pmcanalseg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 98    """Download the PMCanalSeg dataset.
 99
100    Args:
101        path: Filepath to a folder where the data is downloaded for further processing.
102        download: Whether to download the data if it is not present.
103
104    Returns:
105        Filepath where the data is downloaded.
106    """
107    os.makedirs(path, exist_ok=True)
108    manifest = _get_manifest(path)
109
110    missing = [entry for entry in manifest if not os.path.exists(os.path.join(path, entry["directory"], entry["filename"]))]  # noqa
111    if missing and not download:
112        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
113
114    for entry in tqdm(missing, desc="Downloading PMCanalSeg"):
115        out_dir = os.path.join(path, entry["directory"])
116        os.makedirs(out_dir, exist_ok=True)
117        out_path = os.path.join(out_dir, entry["filename"])
118        url = f"{BASE_URL}/api/access/datafile/{entry['id']}"
119        _download_file(url, out_path, entry["md5"])
120
121    return path

Download the PMCanalSeg 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_pmcanalseg_paths( path: Union[os.PathLike, str], label_choice: Literal['mandibular', 'pterygopalatine'] = 'mandibular', download: bool = False) -> Tuple[List[str], List[str]]:
124def get_pmcanalseg_paths(
125    path: Union[os.PathLike, str],
126    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
127    download: bool = False,
128) -> Tuple[List[str], List[str]]:
129    """Get paths to the PMCanalSeg data.
130
131    Args:
132        path: Filepath to a folder where the data is downloaded for further processing.
133        label_choice: The choice of canal to segment. Either 'mandibular' (from the 'lower'
134            mandible CBCT scans) or 'pterygopalatine' (from the 'upper' maxillary CBCT scans).
135        download: Whether to download the data if it is not present.
136
137    Returns:
138        List of filepaths for the image data.
139        List of filepaths for the label data.
140    """
141    if label_choice not in ("mandibular", "pterygopalatine"):
142        raise ValueError(f"'{label_choice}' is not a valid label choice. Please choose 'mandibular' or 'pterygopalatine'.")  # noqa
143
144    data_dir = get_pmcanalseg_data(path, download)
145    subdir = "lower" if label_choice == "mandibular" else "upper"
146
147    image_paths = natsorted(glob(os.path.join(data_dir, subdir, "Patient_*", "image.nii.gz")))
148    gt_paths = [p.replace("image.nii.gz", "label.nii.gz") for p in image_paths]
149
150    image_paths = [p for p, g in zip(image_paths, gt_paths) if os.path.exists(g)]
151    gt_paths = [g for g in gt_paths if os.path.exists(g)]
152
153    return image_paths, gt_paths

Get paths to the PMCanalSeg data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • label_choice: The choice of canal to segment. Either 'mandibular' (from the 'lower' mandible CBCT scans) or 'pterygopalatine' (from the 'upper' maxillary CBCT scans).
  • 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_pmcanalseg_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], label_choice: Literal['mandibular', 'pterygopalatine'] = 'mandibular', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
156def get_pmcanalseg_dataset(
157    path: Union[os.PathLike, str],
158    patch_shape: Tuple[int, ...],
159    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
160    resize_inputs: bool = False,
161    download: bool = False,
162    **kwargs
163) -> Dataset:
164    """Get the PMCanalSeg dataset for segmentation of the mandibular or pterygopalatine canal in CBCT.
165
166    Args:
167        path: Filepath to a folder where the data is downloaded for further processing.
168        patch_shape: The patch shape to use for training.
169        label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
170        resize_inputs: Whether to resize the inputs to the patch shape.
171        download: Whether to download the data if it is not present.
172        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
173
174    Returns:
175        The segmentation dataset.
176    """
177    image_paths, gt_paths = get_pmcanalseg_paths(path, label_choice, download)
178
179    if resize_inputs:
180        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
181        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
182            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
183        )
184
185    return torch_em.default_segmentation_dataset(
186        raw_paths=image_paths,
187        raw_key="data",
188        label_paths=gt_paths,
189        label_key="data",
190        patch_shape=patch_shape,
191        is_seg_dataset=True,
192        **kwargs
193    )

Get the PMCanalSeg dataset for segmentation of the mandibular or pterygopalatine canal in CBCT.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
  • 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_pmcanalseg_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], label_choice: Literal['mandibular', 'pterygopalatine'] = 'mandibular', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
196def get_pmcanalseg_loader(
197    path: Union[os.PathLike, str],
198    batch_size: int,
199    patch_shape: Tuple[int, ...],
200    label_choice: Literal["mandibular", "pterygopalatine"] = "mandibular",
201    resize_inputs: bool = False,
202    download: bool = False,
203    **kwargs
204) -> DataLoader:
205    """Get the PMCanalSeg dataloader for segmentation of the mandibular or pterygopalatine canal in CBCT.
206
207    Args:
208        path: Filepath to a folder where the data is downloaded for further processing.
209        batch_size: The batch size for training.
210        patch_shape: The patch shape to use for training.
211        label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
212        resize_inputs: Whether to resize the inputs to the patch shape.
213        download: Whether to download the data if it is not present.
214        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
215
216    Returns:
217        The DataLoader.
218    """
219    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
220    dataset = get_pmcanalseg_dataset(path, patch_shape, label_choice, resize_inputs, download, **ds_kwargs)
221    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the PMCanalSeg dataloader for segmentation of the mandibular or pterygopalatine canal in CBCT.

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.
  • label_choice: The choice of canal to segment. Either 'mandibular' or 'pterygopalatine'.
  • 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.