torch_em.data.datasets.medical.totalsegmentator

The TotalSegmentator dataset contains annotations for 117 anatomical structures in CT scans.

Two versions of the dataset can be selected via the 'version' argument:

Both versions ship an official split in 'meta.csv' ('v2': train / val / test, 'v3': train / test only, i.e. no 'val' split is available for 'v3') and provide each anatomical structure as a separate binary mask. get_totalsegmentator_data merges these masks into a single semantic label volume per case, where the label id of each structure is its (1-based) position in CLASS_NAMES for 'v2', or CLASS_NAMES_V3 for 'v3' (see CLASS_IDS for the 'v2' name -> id mapping). This is the class order of the 'total' task in the TotalSegmentator repository (https://github.com/wasserth/TotalSegmentator). 'v3' relabels 'vertebrae_S1' as 'vertebrae_L6' as part of its corrected vertebrae labeling (verified on the released 'v3' archive), all other 116 structures are unchanged between 'v2' and 'v3'. The masks of a few structures may overlap, in this case the structure with the higher label id takes precedence.

This dataset is from the publication https://doi.org/10.1148/ryai.230024. Please cite it if you use this dataset in your research.

  1"""The TotalSegmentator dataset contains annotations for 117 anatomical structures in CT scans.
  2
  3Two versions of the dataset can be selected via the 'version' argument:
  4- 'v2' (default, for backward compatibility): 1228 CT volumes, located at https://doi.org/10.5281/zenodo.10047292.
  5- 'v3': 1939 CT volumes (291 additional pediatric CTs on top of the 1228 from 'v2', plus corrected and refined
  6  labels, especially for bones and vertebrae), located at https://doi.org/10.5281/zenodo.22688904.
  7
  8Both versions ship an official split in 'meta.csv' ('v2': train / val / test, 'v3': train / test only, i.e. no
  9'val' split is available for 'v3') and provide each anatomical structure as a separate binary mask.
 10`get_totalsegmentator_data` merges these masks into a single semantic label volume per case, where the label id
 11of each structure is its (1-based) position in `CLASS_NAMES` for 'v2', or `CLASS_NAMES_V3` for 'v3' (see
 12`CLASS_IDS` for the 'v2' name -> id mapping). This is the class order of the 'total' task in the
 13TotalSegmentator repository (https://github.com/wasserth/TotalSegmentator). 'v3' relabels 'vertebrae_S1' as
 14'vertebrae_L6' as part of its corrected vertebrae labeling (verified on the released 'v3' archive), all other
 15116 structures are unchanged between 'v2' and 'v3'. The masks of a few structures may overlap, in this case the
 16structure with the higher label id takes precedence.
 17
 18This dataset is from the publication https://doi.org/10.1148/ryai.230024.
 19Please cite it if you use this dataset in your research.
 20"""
 21
 22import os
 23from glob import glob
 24from concurrent import futures
 25from typing import Union, Tuple, Literal, List, Optional
 26
 27import numpy as np
 28from tqdm import tqdm
 29
 30from torch.utils.data import Dataset, DataLoader
 31
 32import torch_em
 33
 34from .. import util
 35
 36
 37URLS = {
 38    "v2": "https://zenodo.org/records/10047292/files/Totalsegmentator_dataset_v201.zip",
 39    "v3": "https://zenodo.org/records/22688904/files/Totalsegmentator_dataset_v300.zip",
 40}
 41CHECKSUMS = {
 42    "v2": "741dbc911a768e2ac2671c66d55332f7302ad624c915a57d08b142d8bdf0ca26",
 43    "v3": "b56ae18553853ff256fb0eef3a02e322fbc6c862d9015238c8c46b1b76fa027b",
 44}
 45DATA_DIRNAMES = {
 46    "v2": "Totalsegmentator_dataset_v201",
 47    "v3": "Totalsegmentator_dataset_v300",
 48}
 49# The 'v2' archive has no top-level folder (files extract directly into the data folder), while the
 50# 'v3' archive already contains a top-level folder matching 'DATA_DIRNAMES["v3"]'.
 51HAS_TOP_LEVEL_DIR = {"v2": False, "v3": True}
 52
 53CLASS_NAMES = [
 54    "spleen", "kidney_right", "kidney_left", "gallbladder", "liver", "stomach", "pancreas", "adrenal_gland_right",
 55    "adrenal_gland_left", "lung_upper_lobe_left", "lung_lower_lobe_left", "lung_upper_lobe_right",
 56    "lung_middle_lobe_right", "lung_lower_lobe_right", "esophagus", "trachea", "thyroid_gland", "small_bowel",
 57    "duodenum", "colon", "urinary_bladder", "prostate", "kidney_cyst_left", "kidney_cyst_right", "sacrum",
 58    "vertebrae_S1", "vertebrae_L5", "vertebrae_L4", "vertebrae_L3", "vertebrae_L2", "vertebrae_L1", "vertebrae_T12",
 59    "vertebrae_T11", "vertebrae_T10", "vertebrae_T9", "vertebrae_T8", "vertebrae_T7", "vertebrae_T6", "vertebrae_T5",
 60    "vertebrae_T4", "vertebrae_T3", "vertebrae_T2", "vertebrae_T1", "vertebrae_C7", "vertebrae_C6", "vertebrae_C5",
 61    "vertebrae_C4", "vertebrae_C3", "vertebrae_C2", "vertebrae_C1", "heart", "aorta", "pulmonary_vein",
 62    "brachiocephalic_trunk", "subclavian_artery_right", "subclavian_artery_left", "common_carotid_artery_right",
 63    "common_carotid_artery_left", "brachiocephalic_vein_left", "brachiocephalic_vein_right", "atrial_appendage_left",
 64    "superior_vena_cava", "inferior_vena_cava", "portal_vein_and_splenic_vein", "iliac_artery_left",
 65    "iliac_artery_right", "iliac_vena_left", "iliac_vena_right", "humerus_left", "humerus_right", "scapula_left",
 66    "scapula_right", "clavicula_left", "clavicula_right", "femur_left", "femur_right", "hip_left", "hip_right",
 67    "spinal_cord", "gluteus_maximus_left", "gluteus_maximus_right", "gluteus_medius_left", "gluteus_medius_right",
 68    "gluteus_minimus_left", "gluteus_minimus_right", "autochthon_left", "autochthon_right", "iliopsoas_left",
 69    "iliopsoas_right", "brain", "skull", "rib_left_1", "rib_left_2", "rib_left_3", "rib_left_4", "rib_left_5",
 70    "rib_left_6", "rib_left_7", "rib_left_8", "rib_left_9", "rib_left_10", "rib_left_11", "rib_left_12", "rib_right_1",
 71    "rib_right_2", "rib_right_3", "rib_right_4", "rib_right_5", "rib_right_6", "rib_right_7", "rib_right_8",
 72    "rib_right_9", "rib_right_10", "rib_right_11", "rib_right_12", "sternum", "costal_cartilages",
 73]
 74"""The anatomical structures of the TotalSegmentator CT dataset. The label id of a structure is its 1-based index."""
 75
 76CLASS_IDS = {name: i + 1 for i, name in enumerate(CLASS_NAMES)}
 77"""Mapping from the name of an anatomical structure to its label id in the merged label volumes."""
 78
 79CLASS_NAMES_V3 = [name if name != "vertebrae_S1" else "vertebrae_L6" for name in CLASS_NAMES]
 80"""The anatomical structures of the 'v3' TotalSegmentator CT dataset. Identical to `CLASS_NAMES`, except that
 81'vertebrae_S1' was replaced by 'vertebrae_L6' as part of the corrected vertebrae labeling in 'v3' (verified on
 82the per-case 'segmentations' folders of the released 'v3' archive, all 117 structures are otherwise unchanged).
 83"""
 84
 85CLASS_NAMES_BY_VERSION = {"v2": CLASS_NAMES, "v3": CLASS_NAMES_V3}
 86
 87
 88def merge_segmentations(case_dir: str, class_names: List[str], label_name: str = "labels.nii.gz") -> str:
 89    """Merge the per-class binary masks of one TotalSegmentator case into a single semantic label volume.
 90
 91    The merged volume is stored as nifti next to the image. If it already exists it is not recomputed,
 92    so that a partially finished conversion can be resumed.
 93
 94    Args:
 95        case_dir: The folder of the case, which contains the 'segmentations' sub-folder.
 96        class_names: The class names in label id order (the first class gets id 1).
 97        label_name: The filename of the merged label volume.
 98
 99    Returns:
100        The filepath to the merged label volume.
101    """
102    import nibabel as nib
103
104    label_path = os.path.join(case_dir, label_name)
105    if os.path.exists(label_path):
106        return label_path
107
108    labels, affine, header = None, None, None
109    for class_id, class_name in enumerate(class_names, start=1):
110        mask_nii = nib.load(os.path.join(case_dir, "segmentations", f"{class_name}.nii.gz"))
111        mask = np.asarray(mask_nii.dataobj) > 0
112        if labels is None:
113            labels = np.zeros(mask.shape, dtype="uint8")
114            affine, header = mask_nii.affine, mask_nii.header
115        labels[mask] = class_id
116
117    # Write to a temporary path first, so that an interrupted conversion is not mistaken for a complete one.
118    tmp_path = os.path.join(case_dir, f"{label_name}.incomplete.nii.gz")
119    nib.save(nib.Nifti1Image(labels, affine, header), tmp_path)
120    os.replace(tmp_path, label_path)
121    return label_path
122
123
124def merge_all_segmentations(case_dirs: List[str], class_names: List[str], n_workers: Optional[int] = None) -> None:
125    """Merge the per-class binary masks of all cases into semantic label volumes.
126
127    Args:
128        case_dirs: The case folders to process.
129        class_names: The class names in label id order.
130        n_workers: The number of parallel workers. By default the number of CPUs (at most 16) is used.
131    """
132    if all(os.path.exists(os.path.join(case_dir, "labels.nii.gz")) for case_dir in case_dirs):
133        return
134
135    if n_workers is None:
136        n_cpus = len(os.sched_getaffinity(0)) if hasattr(os, "sched_getaffinity") else (os.cpu_count() or 1)
137        n_workers = min(16, n_cpus)
138
139    with futures.ProcessPoolExecutor(n_workers) as pool:
140        tasks = [pool.submit(merge_segmentations, case_dir, class_names) for case_dir in case_dirs]
141        for task in tqdm(futures.as_completed(tasks), total=len(tasks), desc="Merge the per-class segmentations"):
142            task.result()
143
144
145def read_split(meta_csv: str, split: str, valid_splits: Tuple[str, ...] = ("train", "val", "test")) -> List[str]:
146    """Read the case ids of a split from the TotalSegmentator 'meta.csv'.
147
148    Args:
149        meta_csv: The path to the 'meta.csv' file.
150        split: The choice of data split.
151        valid_splits: The splits available in this dataset.
152
153    Returns:
154        The case ids of the split.
155    """
156    import pandas as pd
157
158    if split not in valid_splits:
159        raise ValueError(f"'{split}' is not a valid split. Choose one of {valid_splits}.")
160
161    meta = pd.read_csv(meta_csv, sep=";", encoding="utf-8-sig")
162    return sorted(meta[meta["split"] == split]["image_id"].tolist())
163
164
165def get_totalsegmentator_data(
166    path: Union[os.PathLike, str],
167    download: bool = False,
168    n_workers: Optional[int] = None,
169    version: Literal["v2", "v3"] = "v2",
170) -> str:
171    """Download the TotalSegmentator CT dataset and merge the per-class masks into semantic label volumes.
172
173    Args:
174        path: Filepath to a folder where the data is downloaded for further processing.
175        download: Whether to download the data if it is not present.
176        n_workers: The number of parallel workers for merging the per-class masks.
177        version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs, incl. pediatric CTs
178            and refined bone / vertebrae labels).
179
180    Returns:
181        Filepath where the data is downloaded.
182    """
183    if version not in URLS:
184        raise ValueError(f"'{version}' is not a valid version. Choose one of {list(URLS.keys())}.")
185
186    dirname = DATA_DIRNAMES[version]
187    data_dir = os.path.join(path, dirname)
188    if not os.path.exists(os.path.join(data_dir, "meta.csv")):
189        os.makedirs(path, exist_ok=True)
190        zip_path = os.path.join(path, f"{dirname}.zip")
191        util.download_source(path=zip_path, url=URLS[version], download=download, checksum=CHECKSUMS[version])
192        # Extract into 'data_dir' if the archive has no top-level folder, otherwise into 'path'
193        # (the archive's own top-level folder then becomes 'data_dir').
194        util.unzip(zip_path=zip_path, dst=path if HAS_TOP_LEVEL_DIR[version] else data_dir)
195
196    case_dirs = sorted(glob(os.path.join(data_dir, "s*")))
197    merge_all_segmentations(case_dirs, CLASS_NAMES_BY_VERSION[version], n_workers)
198
199    return data_dir
200
201
202def get_totalsegmentator_paths(
203    path: Union[os.PathLike, str],
204    split: Literal['train', 'val', 'test'],
205    download: bool = False,
206    version: Literal["v2", "v3"] = "v2",
207) -> Tuple[List[str], List[str]]:
208    """Get paths to the TotalSegmentator CT data.
209
210    Args:
211        path: Filepath to a folder where the data is downloaded for further processing.
212        split: The choice of data split.
213        download: Whether to download the data if it is not present.
214        version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs).
215
216    Returns:
217        List of filepaths for the image data.
218        List of filepaths for the label data.
219    """
220    data_dir = get_totalsegmentator_data(path, download, version=version)
221    case_ids = read_split(os.path.join(data_dir, "meta.csv"), split)
222
223    raw_paths = [os.path.join(data_dir, case_id, "ct.nii.gz") for case_id in case_ids]
224    label_paths = [os.path.join(data_dir, case_id, "labels.nii.gz") for case_id in case_ids]
225    assert all(os.path.exists(p) for p in raw_paths + label_paths)
226
227    return raw_paths, label_paths
228
229
230def get_totalsegmentator_dataset(
231    path: Union[os.PathLike, str],
232    patch_shape: Tuple[int, ...],
233    split: Literal['train', 'val', 'test'],
234    resize_inputs: bool = False,
235    download: bool = False,
236    version: Literal["v2", "v3"] = "v2",
237    **kwargs
238) -> Dataset:
239    """Get the TotalSegmentator dataset for segmentation of anatomical structures in CT.
240
241    Args:
242        path: Filepath to a folder where the data is downloaded for further processing.
243        patch_shape: The patch shape to use for training.
244        split: The choice of data split.
245        resize_inputs: Whether to resize inputs to the desired patch shape.
246        download: Whether to download the data if it is not present.
247        version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
248        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
249
250    Returns:
251        The segmentation dataset.
252    """
253    raw_paths, label_paths = get_totalsegmentator_paths(path, split, download, version=version)
254
255    if resize_inputs:
256        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
257        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
258            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
259        )
260
261    return torch_em.default_segmentation_dataset(
262        raw_paths=raw_paths,
263        raw_key="data",
264        label_paths=label_paths,
265        label_key="data",
266        patch_shape=patch_shape,
267        is_seg_dataset=True,
268        **kwargs
269    )
270
271
272def get_totalsegmentator_loader(
273    path: Union[os.PathLike, str],
274    batch_size: int,
275    patch_shape: Tuple[int, ...],
276    split: Literal['train', 'val', 'test'],
277    resize_inputs: bool = False,
278    download: bool = False,
279    version: Literal["v2", "v3"] = "v2",
280    **kwargs
281) -> DataLoader:
282    """Get the TotalSegmentator dataloader for segmentation of anatomical structures in CT.
283
284    Args:
285        path: Filepath to a folder where the data is downloaded for further processing.
286        batch_size: The batch size for training.
287        patch_shape: The patch shape to use for training.
288        split: The choice of data split.
289        resize_inputs: Whether to resize inputs to the desired patch shape.
290        download: Whether to download the data if it is not present.
291        version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
292        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
293
294    Returns:
295        The DataLoader.
296    """
297    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
298    dataset = get_totalsegmentator_dataset(path, patch_shape, split, resize_inputs, download, version, **ds_kwargs)
299    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URLS = {'v2': 'https://zenodo.org/records/10047292/files/Totalsegmentator_dataset_v201.zip', 'v3': 'https://zenodo.org/records/22688904/files/Totalsegmentator_dataset_v300.zip'}
CHECKSUMS = {'v2': '741dbc911a768e2ac2671c66d55332f7302ad624c915a57d08b142d8bdf0ca26', 'v3': 'b56ae18553853ff256fb0eef3a02e322fbc6c862d9015238c8c46b1b76fa027b'}
DATA_DIRNAMES = {'v2': 'Totalsegmentator_dataset_v201', 'v3': 'Totalsegmentator_dataset_v300'}
HAS_TOP_LEVEL_DIR = {'v2': False, 'v3': True}
CLASS_NAMES = ['spleen', 'kidney_right', 'kidney_left', 'gallbladder', 'liver', 'stomach', 'pancreas', 'adrenal_gland_right', 'adrenal_gland_left', 'lung_upper_lobe_left', 'lung_lower_lobe_left', 'lung_upper_lobe_right', 'lung_middle_lobe_right', 'lung_lower_lobe_right', 'esophagus', 'trachea', 'thyroid_gland', 'small_bowel', 'duodenum', 'colon', 'urinary_bladder', 'prostate', 'kidney_cyst_left', 'kidney_cyst_right', 'sacrum', 'vertebrae_S1', 'vertebrae_L5', 'vertebrae_L4', 'vertebrae_L3', 'vertebrae_L2', 'vertebrae_L1', 'vertebrae_T12', 'vertebrae_T11', 'vertebrae_T10', 'vertebrae_T9', 'vertebrae_T8', 'vertebrae_T7', 'vertebrae_T6', 'vertebrae_T5', 'vertebrae_T4', 'vertebrae_T3', 'vertebrae_T2', 'vertebrae_T1', 'vertebrae_C7', 'vertebrae_C6', 'vertebrae_C5', 'vertebrae_C4', 'vertebrae_C3', 'vertebrae_C2', 'vertebrae_C1', 'heart', 'aorta', 'pulmonary_vein', 'brachiocephalic_trunk', 'subclavian_artery_right', 'subclavian_artery_left', 'common_carotid_artery_right', 'common_carotid_artery_left', 'brachiocephalic_vein_left', 'brachiocephalic_vein_right', 'atrial_appendage_left', 'superior_vena_cava', 'inferior_vena_cava', 'portal_vein_and_splenic_vein', 'iliac_artery_left', 'iliac_artery_right', 'iliac_vena_left', 'iliac_vena_right', 'humerus_left', 'humerus_right', 'scapula_left', 'scapula_right', 'clavicula_left', 'clavicula_right', 'femur_left', 'femur_right', 'hip_left', 'hip_right', 'spinal_cord', 'gluteus_maximus_left', 'gluteus_maximus_right', 'gluteus_medius_left', 'gluteus_medius_right', 'gluteus_minimus_left', 'gluteus_minimus_right', 'autochthon_left', 'autochthon_right', 'iliopsoas_left', 'iliopsoas_right', 'brain', 'skull', 'rib_left_1', 'rib_left_2', 'rib_left_3', 'rib_left_4', 'rib_left_5', 'rib_left_6', 'rib_left_7', 'rib_left_8', 'rib_left_9', 'rib_left_10', 'rib_left_11', 'rib_left_12', 'rib_right_1', 'rib_right_2', 'rib_right_3', 'rib_right_4', 'rib_right_5', 'rib_right_6', 'rib_right_7', 'rib_right_8', 'rib_right_9', 'rib_right_10', 'rib_right_11', 'rib_right_12', 'sternum', 'costal_cartilages']

The anatomical structures of the TotalSegmentator CT dataset. The label id of a structure is its 1-based index.

CLASS_IDS = {'spleen': 1, 'kidney_right': 2, 'kidney_left': 3, 'gallbladder': 4, 'liver': 5, 'stomach': 6, 'pancreas': 7, 'adrenal_gland_right': 8, 'adrenal_gland_left': 9, 'lung_upper_lobe_left': 10, 'lung_lower_lobe_left': 11, 'lung_upper_lobe_right': 12, 'lung_middle_lobe_right': 13, 'lung_lower_lobe_right': 14, 'esophagus': 15, 'trachea': 16, 'thyroid_gland': 17, 'small_bowel': 18, 'duodenum': 19, 'colon': 20, 'urinary_bladder': 21, 'prostate': 22, 'kidney_cyst_left': 23, 'kidney_cyst_right': 24, 'sacrum': 25, 'vertebrae_S1': 26, 'vertebrae_L5': 27, 'vertebrae_L4': 28, 'vertebrae_L3': 29, 'vertebrae_L2': 30, 'vertebrae_L1': 31, 'vertebrae_T12': 32, 'vertebrae_T11': 33, 'vertebrae_T10': 34, 'vertebrae_T9': 35, 'vertebrae_T8': 36, 'vertebrae_T7': 37, 'vertebrae_T6': 38, 'vertebrae_T5': 39, 'vertebrae_T4': 40, 'vertebrae_T3': 41, 'vertebrae_T2': 42, 'vertebrae_T1': 43, 'vertebrae_C7': 44, 'vertebrae_C6': 45, 'vertebrae_C5': 46, 'vertebrae_C4': 47, 'vertebrae_C3': 48, 'vertebrae_C2': 49, 'vertebrae_C1': 50, 'heart': 51, 'aorta': 52, 'pulmonary_vein': 53, 'brachiocephalic_trunk': 54, 'subclavian_artery_right': 55, 'subclavian_artery_left': 56, 'common_carotid_artery_right': 57, 'common_carotid_artery_left': 58, 'brachiocephalic_vein_left': 59, 'brachiocephalic_vein_right': 60, 'atrial_appendage_left': 61, 'superior_vena_cava': 62, 'inferior_vena_cava': 63, 'portal_vein_and_splenic_vein': 64, 'iliac_artery_left': 65, 'iliac_artery_right': 66, 'iliac_vena_left': 67, 'iliac_vena_right': 68, 'humerus_left': 69, 'humerus_right': 70, 'scapula_left': 71, 'scapula_right': 72, 'clavicula_left': 73, 'clavicula_right': 74, 'femur_left': 75, 'femur_right': 76, 'hip_left': 77, 'hip_right': 78, 'spinal_cord': 79, 'gluteus_maximus_left': 80, 'gluteus_maximus_right': 81, 'gluteus_medius_left': 82, 'gluteus_medius_right': 83, 'gluteus_minimus_left': 84, 'gluteus_minimus_right': 85, 'autochthon_left': 86, 'autochthon_right': 87, 'iliopsoas_left': 88, 'iliopsoas_right': 89, 'brain': 90, 'skull': 91, 'rib_left_1': 92, 'rib_left_2': 93, 'rib_left_3': 94, 'rib_left_4': 95, 'rib_left_5': 96, 'rib_left_6': 97, 'rib_left_7': 98, 'rib_left_8': 99, 'rib_left_9': 100, 'rib_left_10': 101, 'rib_left_11': 102, 'rib_left_12': 103, 'rib_right_1': 104, 'rib_right_2': 105, 'rib_right_3': 106, 'rib_right_4': 107, 'rib_right_5': 108, 'rib_right_6': 109, 'rib_right_7': 110, 'rib_right_8': 111, 'rib_right_9': 112, 'rib_right_10': 113, 'rib_right_11': 114, 'rib_right_12': 115, 'sternum': 116, 'costal_cartilages': 117}

Mapping from the name of an anatomical structure to its label id in the merged label volumes.

CLASS_NAMES_V3 = ['spleen', 'kidney_right', 'kidney_left', 'gallbladder', 'liver', 'stomach', 'pancreas', 'adrenal_gland_right', 'adrenal_gland_left', 'lung_upper_lobe_left', 'lung_lower_lobe_left', 'lung_upper_lobe_right', 'lung_middle_lobe_right', 'lung_lower_lobe_right', 'esophagus', 'trachea', 'thyroid_gland', 'small_bowel', 'duodenum', 'colon', 'urinary_bladder', 'prostate', 'kidney_cyst_left', 'kidney_cyst_right', 'sacrum', 'vertebrae_L6', 'vertebrae_L5', 'vertebrae_L4', 'vertebrae_L3', 'vertebrae_L2', 'vertebrae_L1', 'vertebrae_T12', 'vertebrae_T11', 'vertebrae_T10', 'vertebrae_T9', 'vertebrae_T8', 'vertebrae_T7', 'vertebrae_T6', 'vertebrae_T5', 'vertebrae_T4', 'vertebrae_T3', 'vertebrae_T2', 'vertebrae_T1', 'vertebrae_C7', 'vertebrae_C6', 'vertebrae_C5', 'vertebrae_C4', 'vertebrae_C3', 'vertebrae_C2', 'vertebrae_C1', 'heart', 'aorta', 'pulmonary_vein', 'brachiocephalic_trunk', 'subclavian_artery_right', 'subclavian_artery_left', 'common_carotid_artery_right', 'common_carotid_artery_left', 'brachiocephalic_vein_left', 'brachiocephalic_vein_right', 'atrial_appendage_left', 'superior_vena_cava', 'inferior_vena_cava', 'portal_vein_and_splenic_vein', 'iliac_artery_left', 'iliac_artery_right', 'iliac_vena_left', 'iliac_vena_right', 'humerus_left', 'humerus_right', 'scapula_left', 'scapula_right', 'clavicula_left', 'clavicula_right', 'femur_left', 'femur_right', 'hip_left', 'hip_right', 'spinal_cord', 'gluteus_maximus_left', 'gluteus_maximus_right', 'gluteus_medius_left', 'gluteus_medius_right', 'gluteus_minimus_left', 'gluteus_minimus_right', 'autochthon_left', 'autochthon_right', 'iliopsoas_left', 'iliopsoas_right', 'brain', 'skull', 'rib_left_1', 'rib_left_2', 'rib_left_3', 'rib_left_4', 'rib_left_5', 'rib_left_6', 'rib_left_7', 'rib_left_8', 'rib_left_9', 'rib_left_10', 'rib_left_11', 'rib_left_12', 'rib_right_1', 'rib_right_2', 'rib_right_3', 'rib_right_4', 'rib_right_5', 'rib_right_6', 'rib_right_7', 'rib_right_8', 'rib_right_9', 'rib_right_10', 'rib_right_11', 'rib_right_12', 'sternum', 'costal_cartilages']

The anatomical structures of the 'v3' TotalSegmentator CT dataset. Identical to CLASS_NAMES, except that 'vertebrae_S1' was replaced by 'vertebrae_L6' as part of the corrected vertebrae labeling in 'v3' (verified on the per-case 'segmentations' folders of the released 'v3' archive, all 117 structures are otherwise unchanged).

CLASS_NAMES_BY_VERSION = {'v2': ['spleen', 'kidney_right', 'kidney_left', 'gallbladder', 'liver', 'stomach', 'pancreas', 'adrenal_gland_right', 'adrenal_gland_left', 'lung_upper_lobe_left', 'lung_lower_lobe_left', 'lung_upper_lobe_right', 'lung_middle_lobe_right', 'lung_lower_lobe_right', 'esophagus', 'trachea', 'thyroid_gland', 'small_bowel', 'duodenum', 'colon', 'urinary_bladder', 'prostate', 'kidney_cyst_left', 'kidney_cyst_right', 'sacrum', 'vertebrae_S1', 'vertebrae_L5', 'vertebrae_L4', 'vertebrae_L3', 'vertebrae_L2', 'vertebrae_L1', 'vertebrae_T12', 'vertebrae_T11', 'vertebrae_T10', 'vertebrae_T9', 'vertebrae_T8', 'vertebrae_T7', 'vertebrae_T6', 'vertebrae_T5', 'vertebrae_T4', 'vertebrae_T3', 'vertebrae_T2', 'vertebrae_T1', 'vertebrae_C7', 'vertebrae_C6', 'vertebrae_C5', 'vertebrae_C4', 'vertebrae_C3', 'vertebrae_C2', 'vertebrae_C1', 'heart', 'aorta', 'pulmonary_vein', 'brachiocephalic_trunk', 'subclavian_artery_right', 'subclavian_artery_left', 'common_carotid_artery_right', 'common_carotid_artery_left', 'brachiocephalic_vein_left', 'brachiocephalic_vein_right', 'atrial_appendage_left', 'superior_vena_cava', 'inferior_vena_cava', 'portal_vein_and_splenic_vein', 'iliac_artery_left', 'iliac_artery_right', 'iliac_vena_left', 'iliac_vena_right', 'humerus_left', 'humerus_right', 'scapula_left', 'scapula_right', 'clavicula_left', 'clavicula_right', 'femur_left', 'femur_right', 'hip_left', 'hip_right', 'spinal_cord', 'gluteus_maximus_left', 'gluteus_maximus_right', 'gluteus_medius_left', 'gluteus_medius_right', 'gluteus_minimus_left', 'gluteus_minimus_right', 'autochthon_left', 'autochthon_right', 'iliopsoas_left', 'iliopsoas_right', 'brain', 'skull', 'rib_left_1', 'rib_left_2', 'rib_left_3', 'rib_left_4', 'rib_left_5', 'rib_left_6', 'rib_left_7', 'rib_left_8', 'rib_left_9', 'rib_left_10', 'rib_left_11', 'rib_left_12', 'rib_right_1', 'rib_right_2', 'rib_right_3', 'rib_right_4', 'rib_right_5', 'rib_right_6', 'rib_right_7', 'rib_right_8', 'rib_right_9', 'rib_right_10', 'rib_right_11', 'rib_right_12', 'sternum', 'costal_cartilages'], 'v3': ['spleen', 'kidney_right', 'kidney_left', 'gallbladder', 'liver', 'stomach', 'pancreas', 'adrenal_gland_right', 'adrenal_gland_left', 'lung_upper_lobe_left', 'lung_lower_lobe_left', 'lung_upper_lobe_right', 'lung_middle_lobe_right', 'lung_lower_lobe_right', 'esophagus', 'trachea', 'thyroid_gland', 'small_bowel', 'duodenum', 'colon', 'urinary_bladder', 'prostate', 'kidney_cyst_left', 'kidney_cyst_right', 'sacrum', 'vertebrae_L6', 'vertebrae_L5', 'vertebrae_L4', 'vertebrae_L3', 'vertebrae_L2', 'vertebrae_L1', 'vertebrae_T12', 'vertebrae_T11', 'vertebrae_T10', 'vertebrae_T9', 'vertebrae_T8', 'vertebrae_T7', 'vertebrae_T6', 'vertebrae_T5', 'vertebrae_T4', 'vertebrae_T3', 'vertebrae_T2', 'vertebrae_T1', 'vertebrae_C7', 'vertebrae_C6', 'vertebrae_C5', 'vertebrae_C4', 'vertebrae_C3', 'vertebrae_C2', 'vertebrae_C1', 'heart', 'aorta', 'pulmonary_vein', 'brachiocephalic_trunk', 'subclavian_artery_right', 'subclavian_artery_left', 'common_carotid_artery_right', 'common_carotid_artery_left', 'brachiocephalic_vein_left', 'brachiocephalic_vein_right', 'atrial_appendage_left', 'superior_vena_cava', 'inferior_vena_cava', 'portal_vein_and_splenic_vein', 'iliac_artery_left', 'iliac_artery_right', 'iliac_vena_left', 'iliac_vena_right', 'humerus_left', 'humerus_right', 'scapula_left', 'scapula_right', 'clavicula_left', 'clavicula_right', 'femur_left', 'femur_right', 'hip_left', 'hip_right', 'spinal_cord', 'gluteus_maximus_left', 'gluteus_maximus_right', 'gluteus_medius_left', 'gluteus_medius_right', 'gluteus_minimus_left', 'gluteus_minimus_right', 'autochthon_left', 'autochthon_right', 'iliopsoas_left', 'iliopsoas_right', 'brain', 'skull', 'rib_left_1', 'rib_left_2', 'rib_left_3', 'rib_left_4', 'rib_left_5', 'rib_left_6', 'rib_left_7', 'rib_left_8', 'rib_left_9', 'rib_left_10', 'rib_left_11', 'rib_left_12', 'rib_right_1', 'rib_right_2', 'rib_right_3', 'rib_right_4', 'rib_right_5', 'rib_right_6', 'rib_right_7', 'rib_right_8', 'rib_right_9', 'rib_right_10', 'rib_right_11', 'rib_right_12', 'sternum', 'costal_cartilages']}
def merge_segmentations( case_dir: str, class_names: List[str], label_name: str = 'labels.nii.gz') -> str:
 89def merge_segmentations(case_dir: str, class_names: List[str], label_name: str = "labels.nii.gz") -> str:
 90    """Merge the per-class binary masks of one TotalSegmentator case into a single semantic label volume.
 91
 92    The merged volume is stored as nifti next to the image. If it already exists it is not recomputed,
 93    so that a partially finished conversion can be resumed.
 94
 95    Args:
 96        case_dir: The folder of the case, which contains the 'segmentations' sub-folder.
 97        class_names: The class names in label id order (the first class gets id 1).
 98        label_name: The filename of the merged label volume.
 99
100    Returns:
101        The filepath to the merged label volume.
102    """
103    import nibabel as nib
104
105    label_path = os.path.join(case_dir, label_name)
106    if os.path.exists(label_path):
107        return label_path
108
109    labels, affine, header = None, None, None
110    for class_id, class_name in enumerate(class_names, start=1):
111        mask_nii = nib.load(os.path.join(case_dir, "segmentations", f"{class_name}.nii.gz"))
112        mask = np.asarray(mask_nii.dataobj) > 0
113        if labels is None:
114            labels = np.zeros(mask.shape, dtype="uint8")
115            affine, header = mask_nii.affine, mask_nii.header
116        labels[mask] = class_id
117
118    # Write to a temporary path first, so that an interrupted conversion is not mistaken for a complete one.
119    tmp_path = os.path.join(case_dir, f"{label_name}.incomplete.nii.gz")
120    nib.save(nib.Nifti1Image(labels, affine, header), tmp_path)
121    os.replace(tmp_path, label_path)
122    return label_path

Merge the per-class binary masks of one TotalSegmentator case into a single semantic label volume.

The merged volume is stored as nifti next to the image. If it already exists it is not recomputed, so that a partially finished conversion can be resumed.

Arguments:
  • case_dir: The folder of the case, which contains the 'segmentations' sub-folder.
  • class_names: The class names in label id order (the first class gets id 1).
  • label_name: The filename of the merged label volume.
Returns:

The filepath to the merged label volume.

def merge_all_segmentations( case_dirs: List[str], class_names: List[str], n_workers: Optional[int] = None) -> None:
125def merge_all_segmentations(case_dirs: List[str], class_names: List[str], n_workers: Optional[int] = None) -> None:
126    """Merge the per-class binary masks of all cases into semantic label volumes.
127
128    Args:
129        case_dirs: The case folders to process.
130        class_names: The class names in label id order.
131        n_workers: The number of parallel workers. By default the number of CPUs (at most 16) is used.
132    """
133    if all(os.path.exists(os.path.join(case_dir, "labels.nii.gz")) for case_dir in case_dirs):
134        return
135
136    if n_workers is None:
137        n_cpus = len(os.sched_getaffinity(0)) if hasattr(os, "sched_getaffinity") else (os.cpu_count() or 1)
138        n_workers = min(16, n_cpus)
139
140    with futures.ProcessPoolExecutor(n_workers) as pool:
141        tasks = [pool.submit(merge_segmentations, case_dir, class_names) for case_dir in case_dirs]
142        for task in tqdm(futures.as_completed(tasks), total=len(tasks), desc="Merge the per-class segmentations"):
143            task.result()

Merge the per-class binary masks of all cases into semantic label volumes.

Arguments:
  • case_dirs: The case folders to process.
  • class_names: The class names in label id order.
  • n_workers: The number of parallel workers. By default the number of CPUs (at most 16) is used.
def read_split( meta_csv: str, split: str, valid_splits: Tuple[str, ...] = ('train', 'val', 'test')) -> List[str]:
146def read_split(meta_csv: str, split: str, valid_splits: Tuple[str, ...] = ("train", "val", "test")) -> List[str]:
147    """Read the case ids of a split from the TotalSegmentator 'meta.csv'.
148
149    Args:
150        meta_csv: The path to the 'meta.csv' file.
151        split: The choice of data split.
152        valid_splits: The splits available in this dataset.
153
154    Returns:
155        The case ids of the split.
156    """
157    import pandas as pd
158
159    if split not in valid_splits:
160        raise ValueError(f"'{split}' is not a valid split. Choose one of {valid_splits}.")
161
162    meta = pd.read_csv(meta_csv, sep=";", encoding="utf-8-sig")
163    return sorted(meta[meta["split"] == split]["image_id"].tolist())

Read the case ids of a split from the TotalSegmentator 'meta.csv'.

Arguments:
  • meta_csv: The path to the 'meta.csv' file.
  • split: The choice of data split.
  • valid_splits: The splits available in this dataset.
Returns:

The case ids of the split.

def get_totalsegmentator_data( path: Union[os.PathLike, str], download: bool = False, n_workers: Optional[int] = None, version: Literal['v2', 'v3'] = 'v2') -> str:
166def get_totalsegmentator_data(
167    path: Union[os.PathLike, str],
168    download: bool = False,
169    n_workers: Optional[int] = None,
170    version: Literal["v2", "v3"] = "v2",
171) -> str:
172    """Download the TotalSegmentator CT dataset and merge the per-class masks into semantic label volumes.
173
174    Args:
175        path: Filepath to a folder where the data is downloaded for further processing.
176        download: Whether to download the data if it is not present.
177        n_workers: The number of parallel workers for merging the per-class masks.
178        version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs, incl. pediatric CTs
179            and refined bone / vertebrae labels).
180
181    Returns:
182        Filepath where the data is downloaded.
183    """
184    if version not in URLS:
185        raise ValueError(f"'{version}' is not a valid version. Choose one of {list(URLS.keys())}.")
186
187    dirname = DATA_DIRNAMES[version]
188    data_dir = os.path.join(path, dirname)
189    if not os.path.exists(os.path.join(data_dir, "meta.csv")):
190        os.makedirs(path, exist_ok=True)
191        zip_path = os.path.join(path, f"{dirname}.zip")
192        util.download_source(path=zip_path, url=URLS[version], download=download, checksum=CHECKSUMS[version])
193        # Extract into 'data_dir' if the archive has no top-level folder, otherwise into 'path'
194        # (the archive's own top-level folder then becomes 'data_dir').
195        util.unzip(zip_path=zip_path, dst=path if HAS_TOP_LEVEL_DIR[version] else data_dir)
196
197    case_dirs = sorted(glob(os.path.join(data_dir, "s*")))
198    merge_all_segmentations(case_dirs, CLASS_NAMES_BY_VERSION[version], n_workers)
199
200    return data_dir

Download the TotalSegmentator CT dataset and merge the per-class masks into semantic label volumes.

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.
  • n_workers: The number of parallel workers for merging the per-class masks.
  • version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs, incl. pediatric CTs and refined bone / vertebrae labels).
Returns:

Filepath where the data is downloaded.

def get_totalsegmentator_paths( path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False, version: Literal['v2', 'v3'] = 'v2') -> Tuple[List[str], List[str]]:
203def get_totalsegmentator_paths(
204    path: Union[os.PathLike, str],
205    split: Literal['train', 'val', 'test'],
206    download: bool = False,
207    version: Literal["v2", "v3"] = "v2",
208) -> Tuple[List[str], List[str]]:
209    """Get paths to the TotalSegmentator CT data.
210
211    Args:
212        path: Filepath to a folder where the data is downloaded for further processing.
213        split: The choice of data split.
214        download: Whether to download the data if it is not present.
215        version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs).
216
217    Returns:
218        List of filepaths for the image data.
219        List of filepaths for the label data.
220    """
221    data_dir = get_totalsegmentator_data(path, download, version=version)
222    case_ids = read_split(os.path.join(data_dir, "meta.csv"), split)
223
224    raw_paths = [os.path.join(data_dir, case_id, "ct.nii.gz") for case_id in case_ids]
225    label_paths = [os.path.join(data_dir, case_id, "labels.nii.gz") for case_id in case_ids]
226    assert all(os.path.exists(p) for p in raw_paths + label_paths)
227
228    return raw_paths, label_paths

Get paths to the TotalSegmentator CT data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split.
  • download: Whether to download the data if it is not present.
  • version: The version of the dataset. Either 'v2' (1228 CTs) or 'v3' (1939 CTs).
Returns:

List of filepaths for the image data. List of filepaths for the label data.

def get_totalsegmentator_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], split: Literal['train', 'val', 'test'], resize_inputs: bool = False, download: bool = False, version: Literal['v2', 'v3'] = 'v2', **kwargs) -> torch.utils.data.dataset.Dataset:
231def get_totalsegmentator_dataset(
232    path: Union[os.PathLike, str],
233    patch_shape: Tuple[int, ...],
234    split: Literal['train', 'val', 'test'],
235    resize_inputs: bool = False,
236    download: bool = False,
237    version: Literal["v2", "v3"] = "v2",
238    **kwargs
239) -> Dataset:
240    """Get the TotalSegmentator dataset for segmentation of anatomical structures in CT.
241
242    Args:
243        path: Filepath to a folder where the data is downloaded for further processing.
244        patch_shape: The patch shape to use for training.
245        split: The choice of data split.
246        resize_inputs: Whether to resize inputs to the desired patch shape.
247        download: Whether to download the data if it is not present.
248        version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
249        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
250
251    Returns:
252        The segmentation dataset.
253    """
254    raw_paths, label_paths = get_totalsegmentator_paths(path, split, download, version=version)
255
256    if resize_inputs:
257        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
258        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
259            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
260        )
261
262    return torch_em.default_segmentation_dataset(
263        raw_paths=raw_paths,
264        raw_key="data",
265        label_paths=label_paths,
266        label_key="data",
267        patch_shape=patch_shape,
268        is_seg_dataset=True,
269        **kwargs
270    )

Get the TotalSegmentator dataset for segmentation of anatomical structures in CT.

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 choice of data split.
  • resize_inputs: Whether to resize inputs to the desired patch shape.
  • download: Whether to download the data if it is not present.
  • version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
  • kwargs: Additional keyword arguments for torch_em.default_segmentation_dataset.
Returns:

The segmentation dataset.

def get_totalsegmentator_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], split: Literal['train', 'val', 'test'], resize_inputs: bool = False, download: bool = False, version: Literal['v2', 'v3'] = 'v2', **kwargs) -> torch.utils.data.dataloader.DataLoader:
273def get_totalsegmentator_loader(
274    path: Union[os.PathLike, str],
275    batch_size: int,
276    patch_shape: Tuple[int, ...],
277    split: Literal['train', 'val', 'test'],
278    resize_inputs: bool = False,
279    download: bool = False,
280    version: Literal["v2", "v3"] = "v2",
281    **kwargs
282) -> DataLoader:
283    """Get the TotalSegmentator dataloader for segmentation of anatomical structures in CT.
284
285    Args:
286        path: Filepath to a folder where the data is downloaded for further processing.
287        batch_size: The batch size for training.
288        patch_shape: The patch shape to use for training.
289        split: The choice of data split.
290        resize_inputs: Whether to resize inputs to the desired patch shape.
291        download: Whether to download the data if it is not present.
292        version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
293        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
294
295    Returns:
296        The DataLoader.
297    """
298    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
299    dataset = get_totalsegmentator_dataset(path, patch_shape, split, resize_inputs, download, version, **ds_kwargs)
300    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the TotalSegmentator dataloader for segmentation of anatomical structures in CT.

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 choice of data split.
  • resize_inputs: Whether to resize inputs to the desired patch shape.
  • download: Whether to download the data if it is not present.
  • version: The version of the dataset. Either 'v2' (1228 CTs, default) or 'v3' (1939 CTs).
  • kwargs: Additional keyword arguments for torch_em.default_segmentation_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.