torch_em.data.datasets.medical.ms3seg

The MS3SEG dataset contains annotations for three-class segmentation of the ventricles, normal age-related white matter hyperintensities (WMH) and pathological multiple sclerosis (MS) WMH lesions in axial T2-FLAIR brain MRI.

The dataset consists of 100 MS patients acquired on a 1.5T Toshiba scanner, with T1-weighted, T2-weighted, and axial / sagittal T2-FLAIR sequences. Expert annotators delineated three classes on the axial T2-FLAIR images: ventricles, normal WMH and abnormal (pathological) WMH.

NOTE: The data is distributed as password-free RAR archives, which requires the 'p7zip' CLI (or the 'rarfile' python package) to extract.

This dataset is from the publication https://doi.org/10.1038/s41597-026-07184-5. Please cite it if you use this dataset in your research.

  1"""The MS3SEG dataset contains annotations for three-class segmentation of the ventricles, normal
  2age-related white matter hyperintensities (WMH) and pathological multiple sclerosis (MS) WMH lesions
  3in axial T2-FLAIR brain MRI.
  4
  5The dataset consists of 100 MS patients acquired on a 1.5T Toshiba scanner, with T1-weighted,
  6T2-weighted, and axial / sagittal T2-FLAIR sequences. Expert annotators delineated three classes on the
  7axial T2-FLAIR images: ventricles, normal WMH and abnormal (pathological) WMH.
  8
  9NOTE: The data is distributed as password-free RAR archives, which requires the 'p7zip' CLI (or the
 10'rarfile' python package) to extract.
 11
 12This dataset is from the publication https://doi.org/10.1038/s41597-026-07184-5.
 13Please cite it if you use this dataset in your research.
 14"""
 15
 16import os
 17from glob import glob
 18from tqdm import tqdm
 19from natsort import natsorted
 20from typing import Union, Tuple, List
 21
 22import numpy as np
 23
 24from torch.utils.data import Dataset, DataLoader
 25
 26import torch_em
 27
 28from .. import util
 29
 30
 31URLS = {
 32    "nifti_part1": "https://ndownloader.figshare.com/files/61900798",
 33    "nifti_part2": "https://ndownloader.figshare.com/files/61901377",
 34    "nifti_part3": "https://ndownloader.figshare.com/files/61901674",
 35    "masks": "https://ndownloader.figshare.com/files/65733546",
 36}
 37
 38CHECKSUMS = {
 39    "nifti_part1": "6a7499b8a6496b76de13783c43af6194b7e81ac1130137a8db639f51a29a973e",
 40    "nifti_part2": "2637a7971e80527f752edfed9e8674daec454b2f020573e1a6279b895f46c7fe",
 41    "nifti_part3": "9b3c71edb700b424a502741623bed38bd3ebaba7dacaef1c4935efc636661e62",
 42    "masks": "6c5d2fddc5ed89988e8c15f060e1564e9ad5e7a10ed4da45070ace397e5c594c",
 43}
 44
 45LABEL_IDS = {"background": 0, "ventricle": 1, "normal_wmh": 2, "ms_wmh": 3}
 46
 47
 48def _preprocess_inputs(path, nifti_dir, masks_dir, preprocessed_dir):
 49    import h5py
 50    import nibabel as nib
 51
 52    os.makedirs(preprocessed_dir, exist_ok=True)
 53
 54    case_dirs = [p for p in natsorted(glob(os.path.join(nifti_dir, "*"))) if os.path.isdir(p)]
 55    for case_dir in tqdm(case_dirs, desc="Preprocessing the MS3SEG cases"):
 56        case_id = os.path.basename(case_dir)
 57        volume_path = os.path.join(preprocessed_dir, f"{case_id}.h5")
 58        if os.path.exists(volume_path):
 59            continue
 60
 61        raw_path = os.path.join(case_dir, f"{case_id}_FLAIR.nii.gz")
 62        vent_path = os.path.join(masks_dir, "Vent_Masks", case_id, f"{case_id}_Vent_Mask.nii.gz")
 63        nwmh_path = os.path.join(masks_dir, "nWMH_Masks", case_id, f"{case_id}_nWMH_Mask.nii.gz")
 64        abwmh_path = os.path.join(masks_dir, "abWMH_Masks", case_id, f"{case_id}_abWMH_Mask.nii.gz")
 65        if not all(os.path.exists(p) for p in (raw_path, vent_path, nwmh_path, abwmh_path)):
 66            continue
 67
 68        raw = np.asarray(nib.load(raw_path).dataobj)
 69        vent = np.asarray(nib.load(vent_path).dataobj)
 70        nwmh = np.asarray(nib.load(nwmh_path).dataobj)
 71        abwmh = np.asarray(nib.load(abwmh_path).dataobj)
 72
 73        # The abnormal (MS) WMH class takes priority over the normal WMH class, which in turn
 74        # takes priority over the ventricle class, in case of overlapping annotations.
 75        labels = np.zeros(raw.shape, dtype="uint8")
 76        labels[vent > 0] = LABEL_IDS["ventricle"]
 77        labels[nwmh > 0] = LABEL_IDS["normal_wmh"]
 78        labels[abwmh > 0] = LABEL_IDS["ms_wmh"]
 79
 80        with h5py.File(f"{volume_path}.tmp", "w") as f:
 81            f.create_dataset("raw", data=raw, compression="gzip")
 82            f.create_dataset("labels", data=labels, compression="gzip")
 83
 84        os.rename(f"{volume_path}.tmp", volume_path)
 85
 86
 87def get_ms3seg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 88    """Download the MS3SEG dataset.
 89
 90    Args:
 91        path: Filepath to a folder where the data is downloaded for further processing.
 92        download: Whether to download the data if it is not present.
 93
 94    Returns:
 95        Filepath where the preprocessed data is stored.
 96    """
 97    preprocessed_dir = os.path.join(path, "preprocessed")
 98    if os.path.exists(preprocessed_dir) and len(glob(os.path.join(preprocessed_dir, "*.h5"))) > 0:
 99        return preprocessed_dir
100
101    os.makedirs(path, exist_ok=True)
102
103    nifti_dir = os.path.join(path, "MS_100_patient_nifti")
104    masks_dir = os.path.join(path, "MS_100_patient_masks")
105
106    if not os.path.exists(nifti_dir):
107        for name in ("nifti_part1", "nifti_part2", "nifti_part3"):
108            rar_path = os.path.join(path, f"{name}.rar")
109            util.download_source(path=rar_path, url=URLS[name], download=download, checksum=CHECKSUMS[name])
110        util.unzip_rarfile(rar_path=os.path.join(path, "nifti_part1.rar"), dst=path, remove=False)
111
112    if not os.path.exists(masks_dir):
113        rar_path = os.path.join(path, "masks.rar")
114        util.download_source(path=rar_path, url=URLS["masks"], download=download, checksum=CHECKSUMS["masks"])
115        util.unzip_rarfile(rar_path=rar_path, dst=path, remove=False)
116
117    _preprocess_inputs(path, nifti_dir, masks_dir, preprocessed_dir)
118    return preprocessed_dir
119
120
121def get_ms3seg_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
122    """Get paths to the MS3SEG data.
123
124    Args:
125        path: Filepath to a folder where the data is downloaded for further processing.
126        download: Whether to download the data if it is not present.
127
128    Returns:
129        List of filepaths for the hdf5 files, which contain the image data ('raw') and the label data ('labels').
130    """
131    data_dir = get_ms3seg_data(path, download)
132    return natsorted(glob(os.path.join(data_dir, "*.h5")))
133
134
135def get_ms3seg_dataset(
136    path: Union[os.PathLike, str],
137    patch_shape: Tuple[int, ...],
138    resize_inputs: bool = False,
139    download: bool = False,
140    **kwargs
141) -> Dataset:
142    """Get the MS3SEG dataset for three-class segmentation of ventricles, normal WMH and MS lesions.
143
144    Args:
145        path: Filepath to a folder where the data is downloaded for further processing.
146        patch_shape: The patch shape to use for training.
147        resize_inputs: Whether to resize inputs to the desired patch shape.
148        download: Whether to download the data if it is not present.
149        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
150
151    Returns:
152        The segmentation dataset.
153    """
154    volume_paths = get_ms3seg_paths(path, download)
155
156    if resize_inputs:
157        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
158        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
159            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
160        )
161
162    return torch_em.default_segmentation_dataset(
163        raw_paths=volume_paths,
164        raw_key="raw",
165        label_paths=volume_paths,
166        label_key="labels",
167        patch_shape=patch_shape,
168        is_seg_dataset=True,
169        **kwargs
170    )
171
172
173def get_ms3seg_loader(
174    path: Union[os.PathLike, str],
175    batch_size: int,
176    patch_shape: Tuple[int, ...],
177    resize_inputs: bool = False,
178    download: bool = False,
179    **kwargs
180) -> DataLoader:
181    """Get the MS3SEG dataloader for three-class segmentation of ventricles, normal WMH and MS lesions.
182
183    Args:
184        path: Filepath to a folder where the data is downloaded for further processing.
185        batch_size: The batch size for training.
186        patch_shape: The patch shape to use for training.
187        resize_inputs: Whether to resize inputs to the desired patch shape.
188        download: Whether to download the data if it is not present.
189        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
190
191    Returns:
192        The DataLoader.
193    """
194    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
195    dataset = get_ms3seg_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
196    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URLS = {'nifti_part1': 'https://ndownloader.figshare.com/files/61900798', 'nifti_part2': 'https://ndownloader.figshare.com/files/61901377', 'nifti_part3': 'https://ndownloader.figshare.com/files/61901674', 'masks': 'https://ndownloader.figshare.com/files/65733546'}
CHECKSUMS = {'nifti_part1': '6a7499b8a6496b76de13783c43af6194b7e81ac1130137a8db639f51a29a973e', 'nifti_part2': '2637a7971e80527f752edfed9e8674daec454b2f020573e1a6279b895f46c7fe', 'nifti_part3': '9b3c71edb700b424a502741623bed38bd3ebaba7dacaef1c4935efc636661e62', 'masks': '6c5d2fddc5ed89988e8c15f060e1564e9ad5e7a10ed4da45070ace397e5c594c'}
LABEL_IDS = {'background': 0, 'ventricle': 1, 'normal_wmh': 2, 'ms_wmh': 3}
def get_ms3seg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 88def get_ms3seg_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 89    """Download the MS3SEG dataset.
 90
 91    Args:
 92        path: Filepath to a folder where the data is downloaded for further processing.
 93        download: Whether to download the data if it is not present.
 94
 95    Returns:
 96        Filepath where the preprocessed data is stored.
 97    """
 98    preprocessed_dir = os.path.join(path, "preprocessed")
 99    if os.path.exists(preprocessed_dir) and len(glob(os.path.join(preprocessed_dir, "*.h5"))) > 0:
100        return preprocessed_dir
101
102    os.makedirs(path, exist_ok=True)
103
104    nifti_dir = os.path.join(path, "MS_100_patient_nifti")
105    masks_dir = os.path.join(path, "MS_100_patient_masks")
106
107    if not os.path.exists(nifti_dir):
108        for name in ("nifti_part1", "nifti_part2", "nifti_part3"):
109            rar_path = os.path.join(path, f"{name}.rar")
110            util.download_source(path=rar_path, url=URLS[name], download=download, checksum=CHECKSUMS[name])
111        util.unzip_rarfile(rar_path=os.path.join(path, "nifti_part1.rar"), dst=path, remove=False)
112
113    if not os.path.exists(masks_dir):
114        rar_path = os.path.join(path, "masks.rar")
115        util.download_source(path=rar_path, url=URLS["masks"], download=download, checksum=CHECKSUMS["masks"])
116        util.unzip_rarfile(rar_path=rar_path, dst=path, remove=False)
117
118    _preprocess_inputs(path, nifti_dir, masks_dir, preprocessed_dir)
119    return preprocessed_dir

Download the MS3SEG 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 preprocessed data is stored.

def get_ms3seg_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
122def get_ms3seg_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
123    """Get paths to the MS3SEG data.
124
125    Args:
126        path: Filepath to a folder where the data is downloaded for further processing.
127        download: Whether to download the data if it is not present.
128
129    Returns:
130        List of filepaths for the hdf5 files, which contain the image data ('raw') and the label data ('labels').
131    """
132    data_dir = get_ms3seg_data(path, download)
133    return natsorted(glob(os.path.join(data_dir, "*.h5")))

Get paths to the MS3SEG 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 hdf5 files, which contain the image data ('raw') and the label data ('labels').

def get_ms3seg_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
136def get_ms3seg_dataset(
137    path: Union[os.PathLike, str],
138    patch_shape: Tuple[int, ...],
139    resize_inputs: bool = False,
140    download: bool = False,
141    **kwargs
142) -> Dataset:
143    """Get the MS3SEG dataset for three-class segmentation of ventricles, normal WMH and MS lesions.
144
145    Args:
146        path: Filepath to a folder where the data is downloaded for further processing.
147        patch_shape: The patch shape to use for training.
148        resize_inputs: Whether to resize inputs to the desired patch shape.
149        download: Whether to download the data if it is not present.
150        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
151
152    Returns:
153        The segmentation dataset.
154    """
155    volume_paths = get_ms3seg_paths(path, download)
156
157    if resize_inputs:
158        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
159        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
160            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
161        )
162
163    return torch_em.default_segmentation_dataset(
164        raw_paths=volume_paths,
165        raw_key="raw",
166        label_paths=volume_paths,
167        label_key="labels",
168        patch_shape=patch_shape,
169        is_seg_dataset=True,
170        **kwargs
171    )

Get the MS3SEG dataset for three-class segmentation of ventricles, normal WMH and MS lesions.

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.

def get_ms3seg_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
174def get_ms3seg_loader(
175    path: Union[os.PathLike, str],
176    batch_size: int,
177    patch_shape: Tuple[int, ...],
178    resize_inputs: bool = False,
179    download: bool = False,
180    **kwargs
181) -> DataLoader:
182    """Get the MS3SEG dataloader for three-class segmentation of ventricles, normal WMH and MS lesions.
183
184    Args:
185        path: Filepath to a folder where the data is downloaded for further processing.
186        batch_size: The batch size for training.
187        patch_shape: The patch shape to use for training.
188        resize_inputs: Whether to resize inputs to the desired patch shape.
189        download: Whether to download the data if it is not present.
190        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
191
192    Returns:
193        The DataLoader.
194    """
195    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
196    dataset = get_ms3seg_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
197    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the MS3SEG dataloader for three-class segmentation of ventricles, normal WMH and MS lesions.

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_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.