torch_em.data.datasets.medical.mslesseg
MSLesSeg is a dataset for segmentation of multiple sclerosis (MS) lesions in brain MRI.
The dataset consists of 115 longitudinal MRI scans from 75 MS patients, with T1-weighted, T2-weighted and FLAIR sequences, registered to the MNI152 template. Expert-validated lesion segmentation masks are provided for each scan.
The dataset is located at https://doi.org/10.6084/m9.figshare.27919209 (Figshare, CC BY 4.0). The dataset is from the publication https://doi.org/10.1038/s41597-025-05250-y. Please cite it if you use this dataset for your research.
1"""MSLesSeg is a dataset for segmentation of multiple sclerosis (MS) lesions in brain MRI. 2 3The dataset consists of 115 longitudinal MRI scans from 75 MS patients, with T1-weighted, T2-weighted 4and FLAIR sequences, registered to the MNI152 template. Expert-validated lesion segmentation masks are 5provided for each scan. 6 7The dataset is located at https://doi.org/10.6084/m9.figshare.27919209 (Figshare, CC BY 4.0). 8The dataset is from the publication https://doi.org/10.1038/s41597-025-05250-y. 9Please cite it if you use this dataset for your research. 10""" 11 12import os 13from glob import glob 14from natsort import natsorted 15from typing import Union, Tuple, Literal, List 16 17from torch.utils.data import Dataset, DataLoader 18 19import torch_em 20 21from .. import util 22 23 24URL = "https://ndownloader.figshare.com/files/52771814" 25 26MODALITIES = ("FLAIR", "T1", "T2") 27 28 29def get_mslesseg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 30 """Download the MSLesSeg data. 31 32 Args: 33 path: Filepath to a folder where the data is downloaded for further processing. 34 download: Whether to download the data if it is not present. 35 36 Returns: 37 Filepath where the data is downloaded. 38 """ 39 data_dir = os.path.join(path, "MSLesSeg Dataset") 40 if os.path.exists(data_dir): 41 return data_dir 42 43 os.makedirs(path, exist_ok=True) 44 45 zip_path = os.path.join(path, "MSLesSeg_Dataset.zip") 46 util.download_source(path=zip_path, url=URL, download=download) 47 util.unzip(zip_path=zip_path, dst=path) 48 49 return data_dir 50 51 52def get_mslesseg_paths( 53 path: Union[os.PathLike, str], modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", download: bool = False 54) -> Tuple[List[str], List[str]]: 55 """Get paths to the MSLesSeg data. 56 57 Args: 58 path: Filepath to a folder where the data is downloaded for further processing. 59 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 60 download: Whether to download the data if it is not present. 61 62 Returns: 63 List of filepaths for the image data. 64 List of filepaths for the label data. 65 """ 66 data_dir = get_mslesseg_data(path, download) 67 68 if modality not in MODALITIES: 69 raise ValueError(f"'{modality}' is not a valid modality. Choose from {MODALITIES}.") 70 71 # The 'train' split stores scans per timepoint (eg. 'train/P9/T3/P9_T3_MASK.nii.gz'), while the 72 # 'test' split has a single timepoint per patient directory (eg. 'test/P54/P54_MASK.nii.gz'). 73 label_paths = natsorted(glob(os.path.join(data_dir, "*", "P*", "**", "*_MASK.nii.gz"), recursive=True)) 74 75 raw_paths = [] 76 for label_path in label_paths: 77 raw_path = label_path.replace("_MASK.nii.gz", f"_{modality}.nii.gz") 78 assert os.path.exists(raw_path), raw_path 79 raw_paths.append(raw_path) 80 81 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 82 83 return raw_paths, label_paths 84 85 86def get_mslesseg_dataset( 87 path: Union[os.PathLike, str], 88 patch_shape: Tuple[int, ...], 89 modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", 90 resize_inputs: bool = False, 91 download: bool = False, 92 **kwargs 93) -> Dataset: 94 """Get the MSLesSeg dataset for multiple sclerosis lesion segmentation in brain MRI. 95 96 Args: 97 path: Filepath to a folder where the data is downloaded for further processing. 98 patch_shape: The patch shape to use for training. 99 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 100 resize_inputs: Whether to resize inputs to the desired patch shape. 101 download: Whether to download the data if it is not present. 102 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 103 104 Returns: 105 The segmentation dataset. 106 """ 107 raw_paths, label_paths = get_mslesseg_paths(path, modality, download) 108 109 if resize_inputs: 110 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 111 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 112 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 113 ) 114 115 return torch_em.default_segmentation_dataset( 116 raw_paths=raw_paths, 117 raw_key="data", 118 label_paths=label_paths, 119 label_key="data", 120 patch_shape=patch_shape, 121 is_seg_dataset=True, 122 **kwargs 123 ) 124 125 126def get_mslesseg_loader( 127 path: Union[os.PathLike, str], 128 batch_size: int, 129 patch_shape: Tuple[int, ...], 130 modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", 131 resize_inputs: bool = False, 132 download: bool = False, 133 **kwargs 134) -> DataLoader: 135 """Get the MSLesSeg dataloader for multiple sclerosis lesion segmentation in brain MRI. 136 137 Args: 138 path: Filepath to a folder where the data is downloaded for further processing. 139 batch_size: The batch size for training. 140 patch_shape: The patch shape to use for training. 141 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 142 resize_inputs: Whether to resize inputs to the desired patch shape. 143 download: Whether to download the data if it is not present. 144 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 145 146 Returns: 147 The DataLoader. 148 """ 149 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 150 dataset = get_mslesseg_dataset(path, patch_shape, modality, resize_inputs, download, **ds_kwargs) 151 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
30def get_mslesseg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 31 """Download the MSLesSeg data. 32 33 Args: 34 path: Filepath to a folder where the data is downloaded for further processing. 35 download: Whether to download the data if it is not present. 36 37 Returns: 38 Filepath where the data is downloaded. 39 """ 40 data_dir = os.path.join(path, "MSLesSeg Dataset") 41 if os.path.exists(data_dir): 42 return data_dir 43 44 os.makedirs(path, exist_ok=True) 45 46 zip_path = os.path.join(path, "MSLesSeg_Dataset.zip") 47 util.download_source(path=zip_path, url=URL, download=download) 48 util.unzip(zip_path=zip_path, dst=path) 49 50 return data_dir
Download the MSLesSeg 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:
Filepath where the data is downloaded.
53def get_mslesseg_paths( 54 path: Union[os.PathLike, str], modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", download: bool = False 55) -> Tuple[List[str], List[str]]: 56 """Get paths to the MSLesSeg data. 57 58 Args: 59 path: Filepath to a folder where the data is downloaded for further processing. 60 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 61 download: Whether to download the data if it is not present. 62 63 Returns: 64 List of filepaths for the image data. 65 List of filepaths for the label data. 66 """ 67 data_dir = get_mslesseg_data(path, download) 68 69 if modality not in MODALITIES: 70 raise ValueError(f"'{modality}' is not a valid modality. Choose from {MODALITIES}.") 71 72 # The 'train' split stores scans per timepoint (eg. 'train/P9/T3/P9_T3_MASK.nii.gz'), while the 73 # 'test' split has a single timepoint per patient directory (eg. 'test/P54/P54_MASK.nii.gz'). 74 label_paths = natsorted(glob(os.path.join(data_dir, "*", "P*", "**", "*_MASK.nii.gz"), recursive=True)) 75 76 raw_paths = [] 77 for label_path in label_paths: 78 raw_path = label_path.replace("_MASK.nii.gz", f"_{modality}.nii.gz") 79 assert os.path.exists(raw_path), raw_path 80 raw_paths.append(raw_path) 81 82 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 83 84 return raw_paths, label_paths
Get paths to the MSLesSeg data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'.
- 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.
87def get_mslesseg_dataset( 88 path: Union[os.PathLike, str], 89 patch_shape: Tuple[int, ...], 90 modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", 91 resize_inputs: bool = False, 92 download: bool = False, 93 **kwargs 94) -> Dataset: 95 """Get the MSLesSeg dataset for multiple sclerosis lesion segmentation in brain MRI. 96 97 Args: 98 path: Filepath to a folder where the data is downloaded for further processing. 99 patch_shape: The patch shape to use for training. 100 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 101 resize_inputs: Whether to resize inputs to the desired patch shape. 102 download: Whether to download the data if it is not present. 103 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 104 105 Returns: 106 The segmentation dataset. 107 """ 108 raw_paths, label_paths = get_mslesseg_paths(path, modality, download) 109 110 if resize_inputs: 111 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 112 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 113 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 114 ) 115 116 return torch_em.default_segmentation_dataset( 117 raw_paths=raw_paths, 118 raw_key="data", 119 label_paths=label_paths, 120 label_key="data", 121 patch_shape=patch_shape, 122 is_seg_dataset=True, 123 **kwargs 124 )
Get the MSLesSeg dataset for multiple sclerosis lesion segmentation in brain MRI.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- patch_shape: The patch shape to use for training.
- modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'.
- 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.
127def get_mslesseg_loader( 128 path: Union[os.PathLike, str], 129 batch_size: int, 130 patch_shape: Tuple[int, ...], 131 modality: Literal["FLAIR", "T1", "T2"] = "FLAIR", 132 resize_inputs: bool = False, 133 download: bool = False, 134 **kwargs 135) -> DataLoader: 136 """Get the MSLesSeg dataloader for multiple sclerosis lesion segmentation in brain MRI. 137 138 Args: 139 path: Filepath to a folder where the data is downloaded for further processing. 140 batch_size: The batch size for training. 141 patch_shape: The patch shape to use for training. 142 modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'. 143 resize_inputs: Whether to resize inputs to the desired patch shape. 144 download: Whether to download the data if it is not present. 145 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 146 147 Returns: 148 The DataLoader. 149 """ 150 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 151 dataset = get_mslesseg_dataset(path, patch_shape, modality, resize_inputs, download, **ds_kwargs) 152 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the MSLesSeg dataloader for multiple sclerosis lesion segmentation in brain MRI.
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.
- modality: The choice of MRI modality. One of 'FLAIR', 'T1', 'T2'.
- 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.