torch_em.data.datasets.medical.full_head_mri_segmentation

The Full-Head MRI Segmentation dataset contains annotations for whole-head segmentation in T1-weighted MRI, including clinical cases with abnormal brain anatomy.

The dataset consists of 68 anonymized clinical subjects (with aphasia or apraxia after a stroke, plus a few healthy controls, scanned at three institutions) and 4 additional healthy control subjects, each with a manually corrected segmentation of the following 7 tissue classes: background, skin/scalp, skull, CSF, gray matter, white matter and air (air cavities and extracephalic air, not always separated into two classes).

The dataset is located at https://www.kaggle.com/datasets/andrewbirnbaum/full-head-mri-and-segmentation-of-stroke-patients and is distributed under the CC BY-NC-SA 4.0 license.

This dataset is from the publication https://doi.org/10.1117/1.JMI.12.5.054001. Please cite it if you use this dataset in your research.

  1"""The Full-Head MRI Segmentation dataset contains annotations for whole-head segmentation in T1-weighted MRI,
  2including clinical cases with abnormal brain anatomy.
  3
  4The dataset consists of 68 anonymized clinical subjects (with aphasia or apraxia after a stroke, plus a few
  5healthy controls, scanned at three institutions) and 4 additional healthy control subjects, each with a manually
  6corrected segmentation of the following 7 tissue classes: background, skin/scalp, skull, CSF, gray matter,
  7white matter and air (air cavities and extracephalic air, not always separated into two classes).
  8
  9The dataset is located at
 10https://www.kaggle.com/datasets/andrewbirnbaum/full-head-mri-and-segmentation-of-stroke-patients
 11and is distributed under the CC BY-NC-SA 4.0 license.
 12
 13This dataset is from the publication https://doi.org/10.1117/1.JMI.12.5.054001. Please cite it if you use this
 14dataset in your research.
 15"""
 16
 17import os
 18from glob import glob
 19from natsort import natsorted
 20from typing import Union, Tuple, List
 21
 22from torch.utils.data import Dataset, DataLoader
 23
 24import torch_em
 25
 26from .. import util
 27
 28
 29KAGGLE_DATASET = "andrewbirnbaum/full-head-mri-and-segmentation-of-stroke-patients"
 30
 31LABEL_IDS = {
 32    "background": 0, "skin_scalp": 1, "skull": 2, "csf": 3, "gray_matter": 4, "white_matter": 5, "air": 6,
 33}
 34
 35
 36def get_full_head_mri_segmentation_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 37    """Download the Full-Head MRI Segmentation dataset.
 38
 39    Args:
 40        path: Filepath to a folder where the data is downloaded for further processing.
 41        download: Whether to download the data if it is not present.
 42
 43    Returns:
 44        Filepath where the data is downloaded.
 45    """
 46    data_dir = os.path.join(path, "Data")
 47    if os.path.exists(data_dir):
 48        return path
 49
 50    os.makedirs(path, exist_ok=True)
 51    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET, download=download)
 52
 53    zip_paths = glob(os.path.join(path, "*.zip"))
 54    assert len(zip_paths) > 0, f"Could not find the downloaded zip file at '{path}'."
 55    util.unzip(zip_path=zip_paths[0], dst=path)
 56
 57    return path
 58
 59
 60def get_full_head_mri_segmentation_paths(
 61    path: Union[os.PathLike, str], download: bool = False
 62) -> Tuple[List[str], List[str]]:
 63    """Get paths to the Full-Head MRI Segmentation data.
 64
 65    Args:
 66        path: Filepath to a folder where the data is downloaded for further processing.
 67        download: Whether to download the data if it is not present.
 68
 69    Returns:
 70        List of filepaths for the image data.
 71        List of filepaths for the label data.
 72    """
 73    data_dir = get_full_head_mri_segmentation_data(path, download)
 74
 75    raw_paths, label_paths = [], []
 76
 77    # The anonymized clinical subjects (T1-weighted MRI file names end in '_deface.nii').
 78    for raw_path in natsorted(glob(os.path.join(
 79        data_dir, "Data", "Anonymized_Subjects", "T1-Weighted MRI", "*_deface.nii"
 80    ))):
 81        label_path = os.path.join(
 82            data_dir, "Data", "Anonymized_Subjects", "Full-Head Segmentation",
 83            os.path.basename(raw_path).replace("_deface.nii", "_label_deface.nii"),
 84        )
 85        if os.path.exists(label_path):
 86            raw_paths.append(raw_path)
 87            label_paths.append(label_path)
 88
 89    # The healthy control subjects.
 90    for raw_path in natsorted(glob(os.path.join(data_dir, "Data", "Control_Subjects", "T1-Weighted MRI", "*.nii"))):
 91        label_path = os.path.join(
 92            data_dir, "Data", "Control_Subjects", "Full-Head Segmentation",
 93            os.path.basename(raw_path).replace(".nii", "_label.nii"),
 94        )
 95        if os.path.exists(label_path):
 96            raw_paths.append(raw_path)
 97            label_paths.append(label_path)
 98
 99    assert len(raw_paths) > 0 and len(raw_paths) == len(label_paths)
100    return raw_paths, label_paths
101
102
103def get_full_head_mri_segmentation_dataset(
104    path: Union[os.PathLike, str],
105    patch_shape: Tuple[int, ...],
106    resize_inputs: bool = False,
107    download: bool = False,
108    **kwargs
109) -> Dataset:
110    """Get the Full-Head MRI Segmentation dataset for whole-head segmentation.
111
112    Args:
113        path: Filepath to a folder where the data is downloaded for further processing.
114        patch_shape: The patch shape to use for training.
115        resize_inputs: Whether to resize inputs to the desired patch shape.
116        download: Whether to download the data if it is not present.
117        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
118
119    Returns:
120        The segmentation dataset.
121    """
122    raw_paths, label_paths = get_full_head_mri_segmentation_paths(path, download)
123
124    if resize_inputs:
125        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
126        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
127            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
128        )
129
130    return torch_em.default_segmentation_dataset(
131        raw_paths=raw_paths,
132        raw_key="data",
133        label_paths=label_paths,
134        label_key="data",
135        patch_shape=patch_shape,
136        is_seg_dataset=True,
137        **kwargs
138    )
139
140
141def get_full_head_mri_segmentation_loader(
142    path: Union[os.PathLike, str],
143    batch_size: int,
144    patch_shape: Tuple[int, ...],
145    resize_inputs: bool = False,
146    download: bool = False,
147    **kwargs
148) -> DataLoader:
149    """Get the Full-Head MRI Segmentation dataloader for whole-head segmentation.
150
151    Args:
152        path: Filepath to a folder where the data is downloaded for further processing.
153        batch_size: The batch size for training.
154        patch_shape: The patch shape to use for training.
155        resize_inputs: Whether to resize inputs to the desired patch shape.
156        download: Whether to download the data if it is not present.
157        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
158
159    Returns:
160        The DataLoader.
161    """
162    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
163    dataset = get_full_head_mri_segmentation_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
164    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
KAGGLE_DATASET = 'andrewbirnbaum/full-head-mri-and-segmentation-of-stroke-patients'
LABEL_IDS = {'background': 0, 'skin_scalp': 1, 'skull': 2, 'csf': 3, 'gray_matter': 4, 'white_matter': 5, 'air': 6}
def get_full_head_mri_segmentation_data(path: Union[os.PathLike, str], download: bool = False) -> str:
37def get_full_head_mri_segmentation_data(path: Union[os.PathLike, str], download: bool = False) -> str:
38    """Download the Full-Head MRI Segmentation dataset.
39
40    Args:
41        path: Filepath to a folder where the data is downloaded for further processing.
42        download: Whether to download the data if it is not present.
43
44    Returns:
45        Filepath where the data is downloaded.
46    """
47    data_dir = os.path.join(path, "Data")
48    if os.path.exists(data_dir):
49        return path
50
51    os.makedirs(path, exist_ok=True)
52    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET, download=download)
53
54    zip_paths = glob(os.path.join(path, "*.zip"))
55    assert len(zip_paths) > 0, f"Could not find the downloaded zip file at '{path}'."
56    util.unzip(zip_path=zip_paths[0], dst=path)
57
58    return path

Download the Full-Head MRI Segmentation 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_full_head_mri_segmentation_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
 61def get_full_head_mri_segmentation_paths(
 62    path: Union[os.PathLike, str], download: bool = False
 63) -> Tuple[List[str], List[str]]:
 64    """Get paths to the Full-Head MRI Segmentation data.
 65
 66    Args:
 67        path: Filepath to a folder where the data is downloaded for further processing.
 68        download: Whether to download the data if it is not present.
 69
 70    Returns:
 71        List of filepaths for the image data.
 72        List of filepaths for the label data.
 73    """
 74    data_dir = get_full_head_mri_segmentation_data(path, download)
 75
 76    raw_paths, label_paths = [], []
 77
 78    # The anonymized clinical subjects (T1-weighted MRI file names end in '_deface.nii').
 79    for raw_path in natsorted(glob(os.path.join(
 80        data_dir, "Data", "Anonymized_Subjects", "T1-Weighted MRI", "*_deface.nii"
 81    ))):
 82        label_path = os.path.join(
 83            data_dir, "Data", "Anonymized_Subjects", "Full-Head Segmentation",
 84            os.path.basename(raw_path).replace("_deface.nii", "_label_deface.nii"),
 85        )
 86        if os.path.exists(label_path):
 87            raw_paths.append(raw_path)
 88            label_paths.append(label_path)
 89
 90    # The healthy control subjects.
 91    for raw_path in natsorted(glob(os.path.join(data_dir, "Data", "Control_Subjects", "T1-Weighted MRI", "*.nii"))):
 92        label_path = os.path.join(
 93            data_dir, "Data", "Control_Subjects", "Full-Head Segmentation",
 94            os.path.basename(raw_path).replace(".nii", "_label.nii"),
 95        )
 96        if os.path.exists(label_path):
 97            raw_paths.append(raw_path)
 98            label_paths.append(label_path)
 99
100    assert len(raw_paths) > 0 and len(raw_paths) == len(label_paths)
101    return raw_paths, label_paths

Get paths to the Full-Head MRI Segmentation 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_full_head_mri_segmentation_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
104def get_full_head_mri_segmentation_dataset(
105    path: Union[os.PathLike, str],
106    patch_shape: Tuple[int, ...],
107    resize_inputs: bool = False,
108    download: bool = False,
109    **kwargs
110) -> Dataset:
111    """Get the Full-Head MRI Segmentation dataset for whole-head segmentation.
112
113    Args:
114        path: Filepath to a folder where the data is downloaded for further processing.
115        patch_shape: The patch shape to use for training.
116        resize_inputs: Whether to resize inputs to the desired patch shape.
117        download: Whether to download the data if it is not present.
118        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
119
120    Returns:
121        The segmentation dataset.
122    """
123    raw_paths, label_paths = get_full_head_mri_segmentation_paths(path, download)
124
125    if resize_inputs:
126        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
127        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
128            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
129        )
130
131    return torch_em.default_segmentation_dataset(
132        raw_paths=raw_paths,
133        raw_key="data",
134        label_paths=label_paths,
135        label_key="data",
136        patch_shape=patch_shape,
137        is_seg_dataset=True,
138        **kwargs
139    )

Get the Full-Head MRI Segmentation dataset for whole-head segmentation.

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_full_head_mri_segmentation_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:
142def get_full_head_mri_segmentation_loader(
143    path: Union[os.PathLike, str],
144    batch_size: int,
145    patch_shape: Tuple[int, ...],
146    resize_inputs: bool = False,
147    download: bool = False,
148    **kwargs
149) -> DataLoader:
150    """Get the Full-Head MRI Segmentation dataloader for whole-head segmentation.
151
152    Args:
153        path: Filepath to a folder where the data is downloaded for further processing.
154        batch_size: The batch size for training.
155        patch_shape: The patch shape to use for training.
156        resize_inputs: Whether to resize inputs to the desired patch shape.
157        download: Whether to download the data if it is not present.
158        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
159
160    Returns:
161        The DataLoader.
162    """
163    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
164    dataset = get_full_head_mri_segmentation_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
165    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the Full-Head MRI Segmentation dataloader for whole-head segmentation.

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.