torch_em.data.datasets.medical.mswal
MSWAL is the first 3D dataset for multi-class segmentation of whole abdominal lesions in CT.
The dataset consists of 484 publicly released training CT volumes (out of 694 total scans acquired at a single
hospital; the 210 held-out test volumes are not released), annotated for seven lesion classes: gallstones,
kidney stones, liver tumors, kidney tumors, pancreatic cancer, liver cysts and kidney cysts. The label ids are
described in LABEL_IDS.
The dataset is located at https://huggingface.co/datasets/zhaodongwu/MSWAL (openly accessible, no gating). This dataset is from the publication https://doi.org/10.1007/978-3-032-04937-7_36 (also on arXiv at https://doi.org/10.48550/arXiv.2503.13560). Please cite it if you use this dataset in your research.
1"""MSWAL is the first 3D dataset for multi-class segmentation of whole abdominal lesions in CT. 2 3The dataset consists of 484 publicly released training CT volumes (out of 694 total scans acquired at a single 4hospital; the 210 held-out test volumes are not released), annotated for seven lesion classes: gallstones, 5kidney stones, liver tumors, kidney tumors, pancreatic cancer, liver cysts and kidney cysts. The label ids are 6described in `LABEL_IDS`. 7 8The dataset is located at https://huggingface.co/datasets/zhaodongwu/MSWAL (openly accessible, no gating). 9This dataset is from the publication https://doi.org/10.1007/978-3-032-04937-7_36 (also on arXiv at 10https://doi.org/10.48550/arXiv.2503.13560). Please cite it if you use this dataset in your research. 11""" 12 13import os 14from glob import glob 15from natsort import natsorted 16from typing import Union, Tuple, List 17 18from torch.utils.data import Dataset, DataLoader 19 20import torch_em 21 22from .. import util 23 24 25REPO_ID = "zhaodongwu/MSWAL" 26 27LABEL_IDS = { 28 "background": 0, 29 "gallstone": 1, 30 "kidney_stone": 2, 31 "liver_tumor": 3, 32 "kidney_tumor": 4, 33 "pancreatic_cancer": 5, 34 "liver_cyst": 6, 35 "kidney_cyst": 7, 36} 37"""The mapping of MSWAL label ids to the corresponding lesion classes.""" 38 39 40def get_mswal_data(path: Union[os.PathLike, str], download: bool = False) -> str: 41 """Download the MSWAL dataset. 42 43 Args: 44 path: Filepath to a folder where the data is downloaded for further processing. 45 download: Whether to download the data if it is not present. 46 47 Returns: 48 Filepath where the data is downloaded. 49 """ 50 data_dir = os.path.join(path, "data") 51 if os.path.exists(os.path.join(data_dir, "imagesTr")) and os.path.exists(os.path.join(data_dir, "labelsTr")): 52 return data_dir 53 54 if not download: 55 raise RuntimeError("The dataset is not found and download is set to False.") 56 57 try: 58 from huggingface_hub import snapshot_download 59 except ModuleNotFoundError: 60 raise ModuleNotFoundError( 61 "Please install 'huggingface_hub' to download the MSWAL dataset: 'pip install huggingface_hub'." 62 ) 63 64 os.makedirs(data_dir, exist_ok=True) 65 snapshot_download( 66 repo_id=REPO_ID, repo_type="dataset", local_dir=data_dir, allow_patterns=["imagesTr/*", "labelsTr/*"] 67 ) 68 69 return data_dir 70 71 72def get_mswal_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 73 """Get paths to the MSWAL data. 74 75 Args: 76 path: Filepath to a folder where the data is downloaded for further processing. 77 download: Whether to download the data if it is not present. 78 79 Returns: 80 List of filepaths for the image data. 81 List of filepaths for the label data. 82 """ 83 data_dir = get_mswal_data(path, download) 84 85 raw_paths = natsorted(glob(os.path.join(data_dir, "imagesTr", "*.nii.gz"))) 86 label_paths = natsorted(glob(os.path.join(data_dir, "labelsTr", "*.nii.gz"))) 87 88 if len(raw_paths) == 0 or len(raw_paths) != len(label_paths): 89 raise RuntimeError("Something went wrong with fetching the image and label paths.") 90 91 return raw_paths, label_paths 92 93 94def get_mswal_dataset( 95 path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], resize_inputs: bool = False, 96 download: bool = False, **kwargs 97) -> Dataset: 98 """Get the MSWAL dataset for multi-class whole abdominal lesion segmentation. 99 100 Args: 101 path: Filepath to a folder where the data is downloaded for further processing. 102 patch_shape: The patch shape to use for training. 103 resize_inputs: Whether to resize inputs to the desired patch shape. 104 download: Whether to download the data if it is not present. 105 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 106 107 Returns: 108 The segmentation dataset. 109 """ 110 raw_paths, label_paths = get_mswal_paths(path, download) 111 112 if resize_inputs: 113 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 114 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 115 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 116 ) 117 118 return torch_em.default_segmentation_dataset( 119 raw_paths=raw_paths, 120 raw_key="data", 121 label_paths=label_paths, 122 label_key="data", 123 patch_shape=patch_shape, 124 is_seg_dataset=True, 125 **kwargs 126 ) 127 128 129def get_mswal_loader( 130 path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], resize_inputs: bool = False, 131 download: bool = False, **kwargs 132) -> DataLoader: 133 """Get the MSWAL dataloader for multi-class whole abdominal lesion segmentation. 134 135 Args: 136 path: Filepath to a folder where the data is downloaded for further processing. 137 batch_size: The batch size for training. 138 patch_shape: The patch shape to use for training. 139 resize_inputs: Whether to resize inputs to the desired patch shape. 140 download: Whether to download the data if it is not present. 141 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 142 143 Returns: 144 The DataLoader. 145 """ 146 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 147 dataset = get_mswal_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 148 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
The mapping of MSWAL label ids to the corresponding lesion classes.
41def get_mswal_data(path: Union[os.PathLike, str], download: bool = False) -> str: 42 """Download the MSWAL dataset. 43 44 Args: 45 path: Filepath to a folder where the data is downloaded for further processing. 46 download: Whether to download the data if it is not present. 47 48 Returns: 49 Filepath where the data is downloaded. 50 """ 51 data_dir = os.path.join(path, "data") 52 if os.path.exists(os.path.join(data_dir, "imagesTr")) and os.path.exists(os.path.join(data_dir, "labelsTr")): 53 return data_dir 54 55 if not download: 56 raise RuntimeError("The dataset is not found and download is set to False.") 57 58 try: 59 from huggingface_hub import snapshot_download 60 except ModuleNotFoundError: 61 raise ModuleNotFoundError( 62 "Please install 'huggingface_hub' to download the MSWAL dataset: 'pip install huggingface_hub'." 63 ) 64 65 os.makedirs(data_dir, exist_ok=True) 66 snapshot_download( 67 repo_id=REPO_ID, repo_type="dataset", local_dir=data_dir, allow_patterns=["imagesTr/*", "labelsTr/*"] 68 ) 69 70 return data_dir
Download the MSWAL 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 data is downloaded.
73def get_mswal_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 74 """Get paths to the MSWAL data. 75 76 Args: 77 path: Filepath to a folder where the data is downloaded for further processing. 78 download: Whether to download the data if it is not present. 79 80 Returns: 81 List of filepaths for the image data. 82 List of filepaths for the label data. 83 """ 84 data_dir = get_mswal_data(path, download) 85 86 raw_paths = natsorted(glob(os.path.join(data_dir, "imagesTr", "*.nii.gz"))) 87 label_paths = natsorted(glob(os.path.join(data_dir, "labelsTr", "*.nii.gz"))) 88 89 if len(raw_paths) == 0 or len(raw_paths) != len(label_paths): 90 raise RuntimeError("Something went wrong with fetching the image and label paths.") 91 92 return raw_paths, label_paths
Get paths to the MSWAL 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 image data. List of filepaths for the label data.
95def get_mswal_dataset( 96 path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], resize_inputs: bool = False, 97 download: bool = False, **kwargs 98) -> Dataset: 99 """Get the MSWAL dataset for multi-class whole abdominal lesion segmentation. 100 101 Args: 102 path: Filepath to a folder where the data is downloaded for further processing. 103 patch_shape: The patch shape to use for training. 104 resize_inputs: Whether to resize inputs to the desired patch shape. 105 download: Whether to download the data if it is not present. 106 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 107 108 Returns: 109 The segmentation dataset. 110 """ 111 raw_paths, label_paths = get_mswal_paths(path, download) 112 113 if resize_inputs: 114 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 115 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 116 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 117 ) 118 119 return torch_em.default_segmentation_dataset( 120 raw_paths=raw_paths, 121 raw_key="data", 122 label_paths=label_paths, 123 label_key="data", 124 patch_shape=patch_shape, 125 is_seg_dataset=True, 126 **kwargs 127 )
Get the MSWAL dataset for multi-class whole abdominal lesion segmentation.
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.
130def get_mswal_loader( 131 path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], resize_inputs: bool = False, 132 download: bool = False, **kwargs 133) -> DataLoader: 134 """Get the MSWAL dataloader for multi-class whole abdominal lesion segmentation. 135 136 Args: 137 path: Filepath to a folder where the data is downloaded for further processing. 138 batch_size: The batch size for training. 139 patch_shape: The patch shape to use for training. 140 resize_inputs: Whether to resize inputs to the desired patch shape. 141 download: Whether to download the data if it is not present. 142 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 143 144 Returns: 145 The DataLoader. 146 """ 147 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 148 dataset = get_mswal_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 149 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the MSWAL dataloader for multi-class whole abdominal lesion 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.
- 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.