torch_em.data.datasets.medical.brainptm
The BrainPTM dataset contains annotations for white matter tracts in brain MRI.
The dataset consists of the 60 training cases of the BrainPTM 2021 challenge, with binary masks for up
to four tracts per case: the left and right optic radiation ('OR_left', 'OR_right') and the left and
right corticospinal tract ('CST_left', 'CST_right'). 56 of the 60 cases have all four tracts, the
remaining 4 have only the optic radiation. See also TRACT_NAMES.
NOTE: The T1 image is used, since it shares its grid with the tracts. The release also provides a diffusion-weighted series per case, which is not on the same grid and is not used here.
NOTE: The 15 test cases of the challenge are not used, because their released tract files are documented placeholders rather than the withheld ground truth.
The dataset is located at https://doi.org/10.5281/zenodo.4600679 and is distributed under the UK Non-Commercial Government Licence v2.0. Please cite the Zenodo record if you use this dataset in your research.
1"""The BrainPTM dataset contains annotations for white matter tracts in brain MRI. 2 3The dataset consists of the 60 training cases of the BrainPTM 2021 challenge, with binary masks for up 4to four tracts per case: the left and right optic radiation ('OR_left', 'OR_right') and the left and 5right corticospinal tract ('CST_left', 'CST_right'). 56 of the 60 cases have all four tracts, the 6remaining 4 have only the optic radiation. See also `TRACT_NAMES`. 7 8NOTE: The T1 image is used, since it shares its grid with the tracts. The release also provides a 9diffusion-weighted series per case, which is not on the same grid and is not used here. 10 11NOTE: The 15 test cases of the challenge are not used, because their released tract files are documented 12placeholders rather than the withheld ground truth. 13 14The dataset is located at https://doi.org/10.5281/zenodo.4600679 and is distributed under the UK 15Non-Commercial Government Licence v2.0. 16Please cite the Zenodo record if you use this dataset in your research. 17""" 18 19import os 20from glob import glob 21from natsort import natsorted 22from typing import Union, Tuple, Literal, List 23 24from torch.utils.data import Dataset, DataLoader 25 26import torch_em 27 28from .. import util 29 30 31URLS = { 32 "data": "https://zenodo.org/records/4600679/files/sheba75_data_train.zip?download=1", 33 "tracts": "https://zenodo.org/records/4600679/files/sheba75_tracts_train.zip?download=1", 34} 35 36CHECKSUMS = { 37 "data": "540f3fd72edffe87b06169b796a149429c9e71104341893d3c792d66667ee60e", 38 "tracts": "e524cdaa73746bcb31718956bd4ea9e00aab50b8fe20fd1adc80dbd005ba22fb", 39} 40 41TRACT_NAMES = ["OR_left", "OR_right", "CST_left", "CST_right"] 42"""The white matter tracts of the BrainPTM dataset. Every case has the optic radiation tracts, only 56 43of the 60 cases also have the corticospinal tracts.""" 44 45 46def get_brainptm_data(path: Union[os.PathLike, str], download: bool = False) -> str: 47 """Download the BrainPTM dataset. 48 49 Args: 50 path: Filepath to a folder where the data is downloaded for further processing. 51 download: Whether to download the data if it is not present. 52 53 Returns: 54 Filepath where the data is downloaded. 55 """ 56 data_dir = os.path.join(path, "data_train") 57 tracts_dir = os.path.join(path, "tracts_train") 58 if os.path.exists(data_dir) and os.path.exists(tracts_dir): 59 return path 60 61 os.makedirs(path, exist_ok=True) 62 for name, dst in [("data", data_dir), ("tracts", tracts_dir)]: 63 zip_path = os.path.join(path, f"sheba75_{name}_train.zip") 64 util.download_source(path=zip_path, url=URLS[name], download=download, checksum=CHECKSUMS[name]) 65 util.unzip(zip_path=zip_path, dst=dst, remove=False) 66 67 return path 68 69 70def get_brainptm_paths( 71 path: Union[os.PathLike, str], 72 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 73 download: bool = False, 74) -> Tuple[List[str], List[str]]: 75 """Get paths to the BrainPTM data. 76 77 Args: 78 path: Filepath to a folder where the data is downloaded for further processing. 79 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 80 download: Whether to download the data if it is not present. 81 82 Returns: 83 List of filepaths for the image data. 84 List of filepaths for the label data. 85 """ 86 if tract not in TRACT_NAMES: 87 raise ValueError(f"'{tract}' is not a valid tract. Choose from {TRACT_NAMES}.") 88 89 root = get_brainptm_data(path, download) 90 91 raw_paths, label_paths = [], [] 92 for label_path in natsorted(glob(os.path.join(root, "tracts_train", "case_*", f"{tract}.nii.gz"))): 93 case_id = os.path.basename(os.path.dirname(label_path)) 94 image_path = os.path.join(root, "data_train", case_id, "T1.nii.gz") 95 if os.path.exists(image_path): 96 raw_paths.append(image_path) 97 label_paths.append(label_path) 98 99 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 100 101 return raw_paths, label_paths 102 103 104def get_brainptm_dataset( 105 path: Union[os.PathLike, str], 106 patch_shape: Tuple[int, ...], 107 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 108 resize_inputs: bool = False, 109 download: bool = False, 110 **kwargs 111) -> Dataset: 112 """Get the BrainPTM dataset for white matter tract segmentation. 113 114 Args: 115 path: Filepath to a folder where the data is downloaded for further processing. 116 patch_shape: The patch shape to use for training. 117 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 118 resize_inputs: Whether to resize inputs to the desired patch shape. 119 download: Whether to download the data if it is not present. 120 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 121 122 Returns: 123 The segmentation dataset. 124 """ 125 raw_paths, label_paths = get_brainptm_paths(path, tract, download) 126 127 if resize_inputs: 128 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 129 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 130 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 131 ) 132 133 return torch_em.default_segmentation_dataset( 134 raw_paths=raw_paths, 135 raw_key="data", 136 label_paths=label_paths, 137 label_key="data", 138 patch_shape=patch_shape, 139 is_seg_dataset=True, 140 **kwargs 141 ) 142 143 144def get_brainptm_loader( 145 path: Union[os.PathLike, str], 146 batch_size: int, 147 patch_shape: Tuple[int, ...], 148 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 149 resize_inputs: bool = False, 150 download: bool = False, 151 **kwargs 152) -> DataLoader: 153 """Get the BrainPTM dataloader for white matter tract segmentation. 154 155 Args: 156 path: Filepath to a folder where the data is downloaded for further processing. 157 batch_size: The batch size for training. 158 patch_shape: The patch shape to use for training. 159 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 160 resize_inputs: Whether to resize inputs to the desired patch shape. 161 download: Whether to download the data if it is not present. 162 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 163 164 Returns: 165 The DataLoader. 166 """ 167 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 168 dataset = get_brainptm_dataset(path, patch_shape, tract, resize_inputs, download, **ds_kwargs) 169 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
The white matter tracts of the BrainPTM dataset. Every case has the optic radiation tracts, only 56 of the 60 cases also have the corticospinal tracts.
47def get_brainptm_data(path: Union[os.PathLike, str], download: bool = False) -> str: 48 """Download the BrainPTM dataset. 49 50 Args: 51 path: Filepath to a folder where the data is downloaded for further processing. 52 download: Whether to download the data if it is not present. 53 54 Returns: 55 Filepath where the data is downloaded. 56 """ 57 data_dir = os.path.join(path, "data_train") 58 tracts_dir = os.path.join(path, "tracts_train") 59 if os.path.exists(data_dir) and os.path.exists(tracts_dir): 60 return path 61 62 os.makedirs(path, exist_ok=True) 63 for name, dst in [("data", data_dir), ("tracts", tracts_dir)]: 64 zip_path = os.path.join(path, f"sheba75_{name}_train.zip") 65 util.download_source(path=zip_path, url=URLS[name], download=download, checksum=CHECKSUMS[name]) 66 util.unzip(zip_path=zip_path, dst=dst, remove=False) 67 68 return path
Download the BrainPTM 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.
71def get_brainptm_paths( 72 path: Union[os.PathLike, str], 73 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 74 download: bool = False, 75) -> Tuple[List[str], List[str]]: 76 """Get paths to the BrainPTM data. 77 78 Args: 79 path: Filepath to a folder where the data is downloaded for further processing. 80 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 81 download: Whether to download the data if it is not present. 82 83 Returns: 84 List of filepaths for the image data. 85 List of filepaths for the label data. 86 """ 87 if tract not in TRACT_NAMES: 88 raise ValueError(f"'{tract}' is not a valid tract. Choose from {TRACT_NAMES}.") 89 90 root = get_brainptm_data(path, download) 91 92 raw_paths, label_paths = [], [] 93 for label_path in natsorted(glob(os.path.join(root, "tracts_train", "case_*", f"{tract}.nii.gz"))): 94 case_id = os.path.basename(os.path.dirname(label_path)) 95 image_path = os.path.join(root, "data_train", case_id, "T1.nii.gz") 96 if os.path.exists(image_path): 97 raw_paths.append(image_path) 98 label_paths.append(label_path) 99 100 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 101 102 return raw_paths, label_paths
Get paths to the BrainPTM data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'.
- 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.
105def get_brainptm_dataset( 106 path: Union[os.PathLike, str], 107 patch_shape: Tuple[int, ...], 108 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 109 resize_inputs: bool = False, 110 download: bool = False, 111 **kwargs 112) -> Dataset: 113 """Get the BrainPTM dataset for white matter tract segmentation. 114 115 Args: 116 path: Filepath to a folder where the data is downloaded for further processing. 117 patch_shape: The patch shape to use for training. 118 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 119 resize_inputs: Whether to resize inputs to the desired patch shape. 120 download: Whether to download the data if it is not present. 121 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 122 123 Returns: 124 The segmentation dataset. 125 """ 126 raw_paths, label_paths = get_brainptm_paths(path, tract, download) 127 128 if resize_inputs: 129 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 130 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 131 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 132 ) 133 134 return torch_em.default_segmentation_dataset( 135 raw_paths=raw_paths, 136 raw_key="data", 137 label_paths=label_paths, 138 label_key="data", 139 patch_shape=patch_shape, 140 is_seg_dataset=True, 141 **kwargs 142 )
Get the BrainPTM dataset for white matter tract segmentation.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- patch_shape: The patch shape to use for training.
- tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'.
- 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.
145def get_brainptm_loader( 146 path: Union[os.PathLike, str], 147 batch_size: int, 148 patch_shape: Tuple[int, ...], 149 tract: Literal["OR_left", "OR_right", "CST_left", "CST_right"] = "OR_left", 150 resize_inputs: bool = False, 151 download: bool = False, 152 **kwargs 153) -> DataLoader: 154 """Get the BrainPTM dataloader for white matter tract segmentation. 155 156 Args: 157 path: Filepath to a folder where the data is downloaded for further processing. 158 batch_size: The batch size for training. 159 patch_shape: The patch shape to use for training. 160 tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'. 161 resize_inputs: Whether to resize inputs to the desired patch shape. 162 download: Whether to download the data if it is not present. 163 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 164 165 Returns: 166 The DataLoader. 167 """ 168 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 169 dataset = get_brainptm_dataset(path, patch_shape, tract, resize_inputs, download, **ds_kwargs) 170 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the BrainPTM dataloader for white matter tract 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.
- tract: The choice of white matter tract. One of 'OR_left', 'OR_right', 'CST_left', 'CST_right'.
- 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.