torch_em.data.datasets.medical.totalsegmentator_hip_implant

The TotalSegmentator hip implant dataset contains annotations for hip implants in CT scans.

This is the training dataset for the "hip_implant" task of the TotalSegmentator repository (https://github.com/wasserth/TotalSegmentator), which is distributed separately from the main TotalSegmentator dataset (see torch_em.data.datasets.medical.totalsegmentator). It consists of 71 CT volumes with a single binary label for hip implants (0 = background, 1 = implant). A small number of volumes are negative controls with an entirely empty (all-background) label; these are filtered out by get_totalsegmentator_hip_implant_paths.

The dataset is located at https://doi.org/10.5281/zenodo.20272031 and licensed under CC BY 4.0.

This dataset is part of the TotalSegmentator project, published at https://doi.org/10.1148/ryai.230024. Please cite it if you use this dataset in your research.

  1"""The TotalSegmentator hip implant dataset contains annotations for hip implants in CT scans.
  2
  3This is the training dataset for the "hip_implant" task of the TotalSegmentator repository
  4(https://github.com/wasserth/TotalSegmentator), which is distributed separately from the main
  5TotalSegmentator dataset (see `torch_em.data.datasets.medical.totalsegmentator`). It consists of 71
  6CT volumes with a single binary label for hip implants (0 = background, 1 = implant). A small number
  7of volumes are negative controls with an entirely empty (all-background) label; these are filtered out
  8by `get_totalsegmentator_hip_implant_paths`.
  9
 10The dataset is located at https://doi.org/10.5281/zenodo.20272031 and licensed under CC BY 4.0.
 11
 12This dataset is part of the TotalSegmentator project, published at https://doi.org/10.1148/ryai.230024.
 13Please cite it if you use this dataset in your research.
 14"""
 15
 16import os
 17from glob import glob
 18from typing import Union, Tuple, List
 19
 20from torch.utils.data import Dataset, DataLoader
 21
 22import torch_em
 23
 24from .. import util
 25
 26
 27URL = "https://zenodo.org/records/20272031/files/Dataset260_hip_implant.zip"
 28CHECKSUM = "c5f7d80ca569f2afb4fe0125dce5b74e9cb14e0ae15da9a2526724d1400aec67"
 29
 30LABEL_IDS = {"background": 0, "implant": 1}
 31
 32
 33def get_totalsegmentator_hip_implant_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 34    """Download the TotalSegmentator hip implant dataset.
 35
 36    Args:
 37        path: Filepath to a folder where the data is downloaded for further processing.
 38        download: Whether to download the data if it is not present.
 39
 40    Returns:
 41        Filepath to the folder with the 'imagesTr' and 'labelsTr' folders.
 42    """
 43    # The archive has no top-level folder, hence it is extracted directly into 'path'.
 44    data_dir = path
 45    if os.path.exists(os.path.join(data_dir, "dataset.json")):
 46        return data_dir
 47
 48    os.makedirs(path, exist_ok=True)
 49    zip_path = os.path.join(path, "Dataset260_hip_implant.zip")
 50    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
 51    util.unzip(zip_path=zip_path, dst=data_dir)
 52
 53    return data_dir
 54
 55
 56def get_totalsegmentator_hip_implant_paths(
 57    path: Union[os.PathLike, str], download: bool = False
 58) -> Tuple[List[str], List[str]]:
 59    """Get paths to the TotalSegmentator hip implant data.
 60
 61    Args:
 62        path: Filepath to a folder where the data is downloaded for further processing.
 63        download: Whether to download the data if it is not present.
 64
 65    Returns:
 66        List of filepaths for the image data.
 67        List of filepaths for the label data.
 68    """
 69    import nibabel as nib
 70    import numpy as np
 71
 72    data_dir = get_totalsegmentator_hip_implant_data(path, download)
 73
 74    raw_paths, label_paths = [], []
 75    for raw_path in sorted(glob(os.path.join(data_dir, "imagesTr", "*_0000.nii.gz"))):
 76        case_id = os.path.basename(raw_path)[:-len("_0000.nii.gz")]
 77        label_path = os.path.join(data_dir, "labelsTr", f"{case_id}.nii.gz")
 78        assert os.path.exists(label_path), label_path
 79
 80        # Skip the rare negative control case(s), whose label volume is entirely background.
 81        if not np.any(nib.load(label_path).get_fdata()):
 82            continue
 83
 84        raw_paths.append(raw_path)
 85        label_paths.append(label_path)
 86
 87    assert len(raw_paths) > 0
 88    return raw_paths, label_paths
 89
 90
 91def get_totalsegmentator_hip_implant_dataset(
 92    path: Union[os.PathLike, str],
 93    patch_shape: Tuple[int, ...],
 94    resize_inputs: bool = False,
 95    download: bool = False,
 96    **kwargs
 97) -> Dataset:
 98    """Get the TotalSegmentator hip implant dataset for hip implant segmentation in CT.
 99
100    Args:
101        path: Filepath to a folder where the data is downloaded for further processing.
102        patch_shape: The patch shape to use for training.
103        resize_inputs: Whether to resize inputs to the desired patch shape.
104        download: Whether to download the data if it is not present.
105        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
106
107    Returns:
108        The segmentation dataset.
109    """
110    raw_paths, label_paths = get_totalsegmentator_hip_implant_paths(path, download)
111
112    if resize_inputs:
113        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
114        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
115            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
116        )
117
118    return torch_em.default_segmentation_dataset(
119        raw_paths=raw_paths,
120        raw_key="data",
121        label_paths=label_paths,
122        label_key="data",
123        patch_shape=patch_shape,
124        is_seg_dataset=True,
125        **kwargs
126    )
127
128
129def get_totalsegmentator_hip_implant_loader(
130    path: Union[os.PathLike, str],
131    batch_size: int,
132    patch_shape: Tuple[int, ...],
133    resize_inputs: bool = False,
134    download: bool = False,
135    **kwargs
136) -> DataLoader:
137    """Get the TotalSegmentator hip implant dataloader for hip implant segmentation in CT.
138
139    Args:
140        path: Filepath to a folder where the data is downloaded for further processing.
141        batch_size: The batch size for training.
142        patch_shape: The patch shape to use for training.
143        resize_inputs: Whether to resize inputs to the desired patch shape.
144        download: Whether to download the data if it is not present.
145        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
146
147    Returns:
148        The DataLoader.
149    """
150    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
151    dataset = get_totalsegmentator_hip_implant_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
152    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://zenodo.org/records/20272031/files/Dataset260_hip_implant.zip'
CHECKSUM = 'c5f7d80ca569f2afb4fe0125dce5b74e9cb14e0ae15da9a2526724d1400aec67'
LABEL_IDS = {'background': 0, 'implant': 1}
def get_totalsegmentator_hip_implant_data(path: Union[os.PathLike, str], download: bool = False) -> str:
34def get_totalsegmentator_hip_implant_data(path: Union[os.PathLike, str], download: bool = False) -> str:
35    """Download the TotalSegmentator hip implant dataset.
36
37    Args:
38        path: Filepath to a folder where the data is downloaded for further processing.
39        download: Whether to download the data if it is not present.
40
41    Returns:
42        Filepath to the folder with the 'imagesTr' and 'labelsTr' folders.
43    """
44    # The archive has no top-level folder, hence it is extracted directly into 'path'.
45    data_dir = path
46    if os.path.exists(os.path.join(data_dir, "dataset.json")):
47        return data_dir
48
49    os.makedirs(path, exist_ok=True)
50    zip_path = os.path.join(path, "Dataset260_hip_implant.zip")
51    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
52    util.unzip(zip_path=zip_path, dst=data_dir)
53
54    return data_dir

Download the TotalSegmentator hip implant 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 to the folder with the 'imagesTr' and 'labelsTr' folders.

def get_totalsegmentator_hip_implant_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
57def get_totalsegmentator_hip_implant_paths(
58    path: Union[os.PathLike, str], download: bool = False
59) -> Tuple[List[str], List[str]]:
60    """Get paths to the TotalSegmentator hip implant data.
61
62    Args:
63        path: Filepath to a folder where the data is downloaded for further processing.
64        download: Whether to download the data if it is not present.
65
66    Returns:
67        List of filepaths for the image data.
68        List of filepaths for the label data.
69    """
70    import nibabel as nib
71    import numpy as np
72
73    data_dir = get_totalsegmentator_hip_implant_data(path, download)
74
75    raw_paths, label_paths = [], []
76    for raw_path in sorted(glob(os.path.join(data_dir, "imagesTr", "*_0000.nii.gz"))):
77        case_id = os.path.basename(raw_path)[:-len("_0000.nii.gz")]
78        label_path = os.path.join(data_dir, "labelsTr", f"{case_id}.nii.gz")
79        assert os.path.exists(label_path), label_path
80
81        # Skip the rare negative control case(s), whose label volume is entirely background.
82        if not np.any(nib.load(label_path).get_fdata()):
83            continue
84
85        raw_paths.append(raw_path)
86        label_paths.append(label_path)
87
88    assert len(raw_paths) > 0
89    return raw_paths, label_paths

Get paths to the TotalSegmentator hip implant data.

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:

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

def get_totalsegmentator_hip_implant_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
 92def get_totalsegmentator_hip_implant_dataset(
 93    path: Union[os.PathLike, str],
 94    patch_shape: Tuple[int, ...],
 95    resize_inputs: bool = False,
 96    download: bool = False,
 97    **kwargs
 98) -> Dataset:
 99    """Get the TotalSegmentator hip implant dataset for hip implant segmentation in CT.
100
101    Args:
102        path: Filepath to a folder where the data is downloaded for further processing.
103        patch_shape: The patch shape to use for training.
104        resize_inputs: Whether to resize inputs to the desired patch shape.
105        download: Whether to download the data if it is not present.
106        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
107
108    Returns:
109        The segmentation dataset.
110    """
111    raw_paths, label_paths = get_totalsegmentator_hip_implant_paths(path, download)
112
113    if resize_inputs:
114        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
115        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
116            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
117        )
118
119    return torch_em.default_segmentation_dataset(
120        raw_paths=raw_paths,
121        raw_key="data",
122        label_paths=label_paths,
123        label_key="data",
124        patch_shape=patch_shape,
125        is_seg_dataset=True,
126        **kwargs
127    )

Get the TotalSegmentator hip implant dataset for hip implant segmentation 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.
  • resize_inputs: Whether to resize inputs to the desired 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_totalsegmentator_hip_implant_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
130def get_totalsegmentator_hip_implant_loader(
131    path: Union[os.PathLike, str],
132    batch_size: int,
133    patch_shape: Tuple[int, ...],
134    resize_inputs: bool = False,
135    download: bool = False,
136    **kwargs
137) -> DataLoader:
138    """Get the TotalSegmentator hip implant dataloader for hip implant segmentation in CT.
139
140    Args:
141        path: Filepath to a folder where the data is downloaded for further processing.
142        batch_size: The batch size for training.
143        patch_shape: The patch shape to use for training.
144        resize_inputs: Whether to resize inputs to the desired patch shape.
145        download: Whether to download the data if it is not present.
146        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
147
148    Returns:
149        The DataLoader.
150    """
151    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
152    dataset = get_totalsegmentator_hip_implant_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
153    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the TotalSegmentator hip implant dataloader for hip implant segmentation 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.
  • resize_inputs: Whether to resize inputs to the desired 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.