torch_em.data.datasets.medical.mbas
The MBAS dataset contains annotations for multi-class bi-atrial segmentation in late gadolinium-enhanced (LGE) cardiac MRI.
The data was curated for the MBAS 2024 challenge (Multi-class Bi-Atrial Segmentation), which was held together
with the STACOM workshop at MICCAI 2024. The public release consists of 100 3D LGE-MRI, split into the 70
studies of the official training set and the 30 studies of the official additional labeled set, which is
selected with the 'split' argument. Each study comes with a multi-class segmentation mask with the label ids
described in LABEL_IDS: 1 = right atrial wall, 2 = left atrial wall, 3 = right atrial cavity,
4 = left atrial cavity.
The data is located at https://zenodo.org/records/19120533 and is released under a custom license: use for non-commercial AI development and testing is permitted and citation is mandatory, while commercialization or other uses require written permission from the authors.
This dataset is from the publication https://doi.org/10.1016/j.media.2026.104203. Please cite it if you use this dataset in your research.
1"""The MBAS dataset contains annotations for multi-class bi-atrial segmentation in 2late gadolinium-enhanced (LGE) cardiac MRI. 3 4The data was curated for the MBAS 2024 challenge (Multi-class Bi-Atrial Segmentation), which was held together 5with the STACOM workshop at MICCAI 2024. The public release consists of 100 3D LGE-MRI, split into the 70 6studies of the official training set and the 30 studies of the official additional labeled set, which is 7selected with the 'split' argument. Each study comes with a multi-class segmentation mask with the label ids 8described in `LABEL_IDS`: 1 = right atrial wall, 2 = left atrial wall, 3 = right atrial cavity, 94 = left atrial cavity. 10 11The data is located at https://zenodo.org/records/19120533 and is released under a custom license: use for 12non-commercial AI development and testing is permitted and citation is mandatory, while commercialization or 13other uses require written permission from the authors. 14 15This dataset is from the publication https://doi.org/10.1016/j.media.2026.104203. 16Please cite it 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, List, Literal 23 24from torch.utils.data import Dataset, DataLoader 25 26import torch_em 27 28from .. import util 29 30 31URLS = { 32 "train": "https://zenodo.org/records/19120533/files/MBAS_Training_4C.zip?download=1", 33 "val": "https://zenodo.org/records/19120533/files/MBAS_Testing_4C.zip?download=1", 34} 35 36CHECKSUMS = { 37 "train": "a765bff61e642ac82a6aa5fffcc39755f75f753db58c111e6e2f50a88a909816", 38 "val": "27c235e2646b37e55734e3a213ad24c56e87540bb0e0c245b6a935a9d6220883", 39} 40 41LABEL_IDS = {"background": 0, "raw": 1, "law": 2, "ra": 3, "la": 4} 42 43 44def get_mbas_data(path: Union[os.PathLike, str], split: Literal["train", "val"], download: bool = False) -> str: 45 """Download the MBAS dataset. 46 47 Args: 48 path: Filepath to a folder where the data is downloaded for further processing. 49 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 50 download: Whether to download the data if it is not present. 51 52 Returns: 53 Filepath where the data is stored. 54 """ 55 if split not in URLS: 56 raise ValueError(f"'{split}' is not a valid split. Please choose one of {list(URLS.keys())}.") 57 58 data_dir = os.path.join(path, split) 59 if os.path.exists(data_dir): 60 return data_dir 61 62 os.makedirs(path, exist_ok=True) 63 64 zip_path = os.path.join(path, f"MBAS_{split}.zip") 65 util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split]) 66 util.unzip(zip_path=zip_path, dst=data_dir) 67 68 return data_dir 69 70 71def get_mbas_paths( 72 path: Union[os.PathLike, str], split: Literal["train", "val"], download: bool = False 73) -> Tuple[List[str], List[str]]: 74 """Get paths to the MBAS data. 75 76 Args: 77 path: Filepath to a folder where the data is downloaded for further processing. 78 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 79 download: Whether to download the data if it is not present. 80 81 Returns: 82 List of filepaths for the image data. 83 List of filepaths for the label data. 84 """ 85 data_dir = get_mbas_data(path, split, download) 86 87 raw_paths = natsorted(glob(os.path.join(data_dir, "*", "*_image.nii.gz"))) 88 label_paths = natsorted(glob(os.path.join(data_dir, "*", "*_label.nii.gz"))) 89 90 if len(raw_paths) == 0 or len(raw_paths) != len(label_paths): 91 raise RuntimeError("Something went wrong with fetching the image and label paths.") 92 93 return raw_paths, label_paths 94 95 96def get_mbas_dataset( 97 path: Union[os.PathLike, str], 98 patch_shape: Tuple[int, ...], 99 split: Literal["train", "val"], 100 resize_inputs: bool = False, 101 download: bool = False, 102 **kwargs 103) -> Dataset: 104 """Get the MBAS dataset for multi-class bi-atrial segmentation. 105 106 Args: 107 path: Filepath to a folder where the data is downloaded for further processing. 108 patch_shape: The patch shape to use for training. 109 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 110 resize_inputs: Whether to resize inputs to the desired patch shape. 111 download: Whether to download the data if it is not present. 112 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 113 114 Returns: 115 The segmentation dataset. 116 """ 117 raw_paths, label_paths = get_mbas_paths(path, split, download) 118 119 if resize_inputs: 120 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 121 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 122 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 123 ) 124 125 return torch_em.default_segmentation_dataset( 126 raw_paths=raw_paths, 127 raw_key="data", 128 label_paths=label_paths, 129 label_key="data", 130 patch_shape=patch_shape, 131 is_seg_dataset=True, 132 **kwargs 133 ) 134 135 136def get_mbas_loader( 137 path: Union[os.PathLike, str], 138 batch_size: int, 139 patch_shape: Tuple[int, ...], 140 split: Literal["train", "val"], 141 resize_inputs: bool = False, 142 download: bool = False, 143 **kwargs 144) -> DataLoader: 145 """Get the MBAS dataloader for multi-class bi-atrial segmentation. 146 147 Args: 148 path: Filepath to a folder where the data is downloaded for further processing. 149 batch_size: The batch size for training. 150 patch_shape: The patch shape to use for training. 151 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 152 resize_inputs: Whether to resize inputs to the desired patch shape. 153 download: Whether to download the data if it is not present. 154 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 155 156 Returns: 157 The DataLoader. 158 """ 159 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 160 dataset = get_mbas_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 161 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
45def get_mbas_data(path: Union[os.PathLike, str], split: Literal["train", "val"], download: bool = False) -> str: 46 """Download the MBAS dataset. 47 48 Args: 49 path: Filepath to a folder where the data is downloaded for further processing. 50 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 51 download: Whether to download the data if it is not present. 52 53 Returns: 54 Filepath where the data is stored. 55 """ 56 if split not in URLS: 57 raise ValueError(f"'{split}' is not a valid split. Please choose one of {list(URLS.keys())}.") 58 59 data_dir = os.path.join(path, split) 60 if os.path.exists(data_dir): 61 return data_dir 62 63 os.makedirs(path, exist_ok=True) 64 65 zip_path = os.path.join(path, f"MBAS_{split}.zip") 66 util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split]) 67 util.unzip(zip_path=zip_path, dst=data_dir) 68 69 return data_dir
Download the MBAS dataset.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases).
- download: Whether to download the data if it is not present.
Returns:
Filepath where the data is stored.
72def get_mbas_paths( 73 path: Union[os.PathLike, str], split: Literal["train", "val"], download: bool = False 74) -> Tuple[List[str], List[str]]: 75 """Get paths to the MBAS data. 76 77 Args: 78 path: Filepath to a folder where the data is downloaded for further processing. 79 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 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 data_dir = get_mbas_data(path, split, download) 87 88 raw_paths = natsorted(glob(os.path.join(data_dir, "*", "*_image.nii.gz"))) 89 label_paths = natsorted(glob(os.path.join(data_dir, "*", "*_label.nii.gz"))) 90 91 if len(raw_paths) == 0 or len(raw_paths) != len(label_paths): 92 raise RuntimeError("Something went wrong with fetching the image and label paths.") 93 94 return raw_paths, label_paths
Get paths to the MBAS data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases).
- 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.
97def get_mbas_dataset( 98 path: Union[os.PathLike, str], 99 patch_shape: Tuple[int, ...], 100 split: Literal["train", "val"], 101 resize_inputs: bool = False, 102 download: bool = False, 103 **kwargs 104) -> Dataset: 105 """Get the MBAS dataset for multi-class bi-atrial segmentation. 106 107 Args: 108 path: Filepath to a folder where the data is downloaded for further processing. 109 patch_shape: The patch shape to use for training. 110 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 111 resize_inputs: Whether to resize inputs to the desired patch shape. 112 download: Whether to download the data if it is not present. 113 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 114 115 Returns: 116 The segmentation dataset. 117 """ 118 raw_paths, label_paths = get_mbas_paths(path, split, download) 119 120 if resize_inputs: 121 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 122 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 123 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 124 ) 125 126 return torch_em.default_segmentation_dataset( 127 raw_paths=raw_paths, 128 raw_key="data", 129 label_paths=label_paths, 130 label_key="data", 131 patch_shape=patch_shape, 132 is_seg_dataset=True, 133 **kwargs 134 )
Get the MBAS dataset for multi-class bi-atrial segmentation.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- patch_shape: The patch shape to use for training.
- split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases).
- 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.
137def get_mbas_loader( 138 path: Union[os.PathLike, str], 139 batch_size: int, 140 patch_shape: Tuple[int, ...], 141 split: Literal["train", "val"], 142 resize_inputs: bool = False, 143 download: bool = False, 144 **kwargs 145) -> DataLoader: 146 """Get the MBAS dataloader for multi-class bi-atrial segmentation. 147 148 Args: 149 path: Filepath to a folder where the data is downloaded for further processing. 150 batch_size: The batch size for training. 151 patch_shape: The patch shape to use for training. 152 split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases). 153 resize_inputs: Whether to resize inputs to the desired patch shape. 154 download: Whether to download the data if it is not present. 155 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 156 157 Returns: 158 The DataLoader. 159 """ 160 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 161 dataset = get_mbas_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 162 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the MBAS dataloader for multi-class bi-atrial 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.
- split: The choice of data split. Either 'train' (70 cases) or 'val' (30 additional cases).
- 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.