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)
URLS = {'train': 'https://zenodo.org/records/19120533/files/MBAS_Training_4C.zip?download=1', 'val': 'https://zenodo.org/records/19120533/files/MBAS_Testing_4C.zip?download=1'}
CHECKSUMS = {'train': 'a765bff61e642ac82a6aa5fffcc39755f75f753db58c111e6e2f50a88a909816', 'val': '27c235e2646b37e55734e3a213ad24c56e87540bb0e0c245b6a935a9d6220883'}
LABEL_IDS = {'background': 0, 'raw': 1, 'law': 2, 'ra': 3, 'la': 4}
def get_mbas_data( path: Union[os.PathLike, str], split: Literal['train', 'val'], download: bool = False) -> str:
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.

def get_mbas_paths( path: Union[os.PathLike, str], split: Literal['train', 'val'], download: bool = False) -> Tuple[List[str], List[str]]:
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.

def get_mbas_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], split: Literal['train', 'val'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
 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.

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

The DataLoader.