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)
REPO_ID = 'zhaodongwu/MSWAL'
LABEL_IDS = {'background': 0, 'gallstone': 1, 'kidney_stone': 2, 'liver_tumor': 3, 'kidney_tumor': 4, 'pancreatic_cancer': 5, 'liver_cyst': 6, 'kidney_cyst': 7}

The mapping of MSWAL label ids to the corresponding lesion classes.

def get_mswal_data(path: Union[os.PathLike, str], download: bool = False) -> str:
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.

def get_mswal_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
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.

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

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

The DataLoader.