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)
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.
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').
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.