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)
URLS = {'data': 'https://zenodo.org/records/4600679/files/sheba75_data_train.zip?download=1', 'tracts': 'https://zenodo.org/records/4600679/files/sheba75_tracts_train.zip?download=1'}
CHECKSUMS = {'data': '540f3fd72edffe87b06169b796a149429c9e71104341893d3c792d66667ee60e', 'tracts': 'e524cdaa73746bcb31718956bd4ea9e00aab50b8fe20fd1adc80dbd005ba22fb'}
TRACT_NAMES = ['OR_left', 'OR_right', 'CST_left', 'CST_right']

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.

def get_brainptm_data(path: Union[os.PathLike, str], download: bool = False) -> str:
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.

def get_brainptm_paths( path: Union[os.PathLike, str], tract: Literal['OR_left', 'OR_right', 'CST_left', 'CST_right'] = 'OR_left', download: bool = False) -> Tuple[List[str], List[str]]:
 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.

def get_brainptm_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], tract: Literal['OR_left', 'OR_right', 'CST_left', 'CST_right'] = 'OR_left', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
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.

def get_brainptm_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], tract: Literal['OR_left', 'OR_right', 'CST_left', 'CST_right'] = 'OR_left', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
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_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.