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:
- 'v2' (default, for backward compatibility): 1228 CT volumes, located at https://doi.org/10.5281/zenodo.10047292.
- 'v3': 1939 CT volumes (291 additional pediatric CTs on top of the 1228 from 'v2', plus corrected and refined labels, especially for bones and vertebrae), located at https://doi.org/10.5281/zenodo.22688904.
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)
The anatomical structures of the TotalSegmentator CT dataset. The label id of a structure is its 1-based index.
Mapping from the name of an anatomical structure to its label id in the merged label volumes.
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).
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.
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.
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.
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.
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.
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.