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)
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.
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.
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.