torch_em.data.datasets.medical.curvas

The CURVAS dataset contains annotations for pancreas, kidney and liver in abdominal CT scans. Each scan is annotated independently by three raters.

The 'train' split consists of the first 10 of the 20 cases of the official training set, the 'val' split of the 5 cases of the official validation set and the 'test' split of the 65 cases of the official testing set. All cases come with the annotations of all three raters.

This dataset is from the challenge: https://curvas.grand-challenge.org. The dataset is located at: https://zenodo.org/records/13767408, and is from the publication https://doi.org/10.48550/arXiv.2505.08685 Please cite tem if you use this dataset for your research.

  1"""The CURVAS dataset contains annotations for pancreas, kidney and liver
  2in abdominal CT scans. Each scan is annotated independently by three raters.
  3
  4The 'train' split consists of the first 10 of the 20 cases of the official training set, the 'val' split of the
  55 cases of the official validation set and the 'test' split of the 65 cases of the official testing set.
  6All cases come with the annotations of all three raters.
  7
  8This dataset is from the challenge: https://curvas.grand-challenge.org.
  9The dataset is located at: https://zenodo.org/records/13767408,
 10and is from the publication https://doi.org/10.48550/arXiv.2505.08685
 11Please cite tem if you use this dataset for your research.
 12"""
 13
 14import os
 15import shutil
 16import subprocess
 17from tqdm import tqdm
 18from glob import glob
 19from natsort import natsorted
 20from typing import Tuple, Union, Literal, List
 21
 22import numpy as np
 23
 24from torch.utils.data import Dataset, DataLoader
 25
 26import torch_em
 27
 28from .. import util
 29
 30
 31URL = "https://zenodo.org/records/12687192/files/training_set.zip"
 32CHECKSUM = "1126a2205553ae1d4fe5fbaee7ea732aacc4f5a92b96504ed521c23e5a0e3f89"
 33
 34URLS = {
 35    "val": "https://zenodo.org/records/13767408/files/validation_set.zip",
 36    "test": "https://zenodo.org/records/13767408/files/testing_set.zip",
 37}
 38CHECKSUMS = {
 39    "val": "01edfac9a085f06111969821d06c83d164654a6041c2e8ac3b11ed390e7c7028",
 40    "test": "6a70aa241a14184778e25d11cae58b39cbc1d8ca204fea19ce0a34bb5b13f7b5",
 41}
 42H5_DIRS = {"train": "data", "val": "data_val", "test": "data_test"}
 43
 44
 45def _preprocess_data(data_dir, h5_dir):
 46    import h5py
 47    import nibabel as nib
 48
 49    os.makedirs(h5_dir, exist_ok=True)
 50
 51    image_paths = natsorted(glob(os.path.join(data_dir, "*", "image.nii.gz")))
 52    for image_path in tqdm(image_paths, desc="Processing data"):
 53        rater1_path = os.path.join(os.path.dirname(image_path), "annotation_1.nii.gz")
 54        rater2_path = os.path.join(os.path.dirname(image_path), "annotation_2.nii.gz")
 55        rater3_path = os.path.join(os.path.dirname(image_path), "annotation_3.nii.gz")
 56
 57        assert os.path.exists(rater1_path) and os.path.exists(rater2_path) and os.path.exists(rater3_path)
 58
 59        image = nib.load(image_path).get_fdata().astype("float32").transpose(2, 0, 1)
 60
 61        label_r1 = np.rint(nib.load(rater1_path).get_fdata()).astype("uint8").transpose(2, 0, 1)
 62        label_r2 = np.rint(nib.load(rater2_path).get_fdata()).astype("uint8").transpose(2, 0, 1)
 63        label_r3 = np.rint(nib.load(rater3_path).get_fdata()).astype("uint8").transpose(2, 0, 1)
 64
 65        fname = os.path.basename(os.path.dirname(image_path))
 66        chunks = (8, 512, 512)
 67        with h5py.File(os.path.join(h5_dir, f"{fname}.h5"), "w") as f:
 68            f.create_dataset("raw", data=image, compression="gzip", chunks=chunks)
 69            f.create_dataset("labels/rater_1", data=label_r1, compression="gzip", chunks=chunks)
 70            f.create_dataset("labels/rater_2", data=label_r2, compression="gzip", chunks=chunks)
 71            f.create_dataset("labels/rater_3", data=label_r3, compression="gzip", chunks=chunks)
 72
 73    # Remove the nifti files as we don't need them anymore!
 74    shutil.rmtree(data_dir)
 75
 76
 77def get_curvas_data(
 78    path: Union[os.PathLike, str], split: Literal["train", "val", "test"] = "train", download: bool = False
 79) -> str:
 80    """Download the CURVAS dataset.
 81
 82    NOTE: The test split is about 21.6 GB.
 83
 84    Args:
 85        path: Filepath to a folder where the data is downloaded for further processing.
 86        split: The choice of data split.
 87        download: Whether to download the data if it is not present.
 88
 89    Returns:
 90        Filepath where the data is downloaded.
 91    """
 92    if split not in H5_DIRS:
 93        raise ValueError(f"'{split}' is not a valid split. Choose one of {list(H5_DIRS)}.")
 94
 95    data_dir = os.path.join(path, H5_DIRS[split])
 96    if os.path.exists(data_dir):
 97        return data_dir
 98
 99    os.makedirs(path, exist_ok=True)
100
101    if split == "train":
102        zip_path = os.path.join(path, "training_set.zip")
103        util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
104
105        # HACK: The zip file is broken. We fix it using the following script.
106        fixed_zip_path = os.path.join(path, "training_set_fixed.zip")
107        subprocess.run(["zip", "-FF", zip_path, "--out", fixed_zip_path])
108        subprocess.run(["unzip", fixed_zip_path, "-d", path])
109
110        _preprocess_data(os.path.join(path, "training_set"), data_dir)
111
112        # Remove the zip files as we don't need them anymore.
113        os.remove(zip_path)
114        os.remove(fixed_zip_path)
115    else:
116        zip_path = os.path.join(path, os.path.basename(URLS[split]))
117        util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split])
118        util.unzip(zip_path=zip_path, dst=path, remove=False)
119
120        _preprocess_data(os.path.join(path, os.path.splitext(os.path.basename(URLS[split]))[0]), data_dir)
121
122        os.remove(zip_path)
123
124    return data_dir
125
126
127def get_curvas_paths(
128    path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False
129) -> List[str]:
130    """Get paths to the CURVAS data.
131
132    Args:
133        path: Filepath to a folder where the data is downloaded for further processing.
134        split: The choice of data split.
135        download: Whether to download the data if it is not present.
136
137    Returns:
138        List of filepaths for the volumetric data.
139    """
140    data_dir = get_curvas_data(path, split, download)
141    volume_paths = natsorted(glob(os.path.join(data_dir, "*.h5")))
142
143    if split == "train":
144        volume_paths = volume_paths[:10]
145
146    return volume_paths
147
148
149def get_curvas_dataset(
150    path: Union[os.PathLike, str],
151    patch_shape: Tuple[int, ...],
152    split: Literal['train', 'val', 'test'],
153    rater: Literal["1", "2", "3"] = "1",
154    resize_inputs: bool = False,
155    download: bool = False,
156    **kwargs
157) -> Dataset:
158    """Get the CURVAS dataset for pancreas, kidney and liver segmentation.
159
160    Args:
161        path: Filepath to a folder where the data is downloaded for further processing.
162        patch_shape: The patch shape to use for training.
163        split: The choice of data split.
164        rater: The choice of rater providing the annotations.
165        resize_inputs: Whether to resize inputs to the desired patch shape.
166        download: Whether to download the data if it is not present.
167        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
168
169    Returns:
170        The segmentation dataset.
171    """
172    volume_paths = get_curvas_paths(path, split, download)
173
174    if resize_inputs:
175        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
176        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
177            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
178        )
179
180    return torch_em.default_segmentation_dataset(
181        raw_paths=volume_paths,
182        raw_key="raw",
183        label_paths=volume_paths,
184        label_key=f"labels/rater_{rater}",
185        patch_shape=patch_shape,
186        is_seg_dataset=True,
187        **kwargs,
188    )
189
190
191def get_curvas_loader(
192    path: Union[os.PathLike, str],
193    batch_size: int,
194    patch_shape: Tuple[int, ...],
195    split: Literal['train', 'val', 'test'],
196    rater: Literal["1", "2", "3"] = "1",
197    resize_inputs: bool = False,
198    download: bool = False,
199    **kwargs
200) -> DataLoader:
201    """Get the CURVAS dataloader for pancreas, kidney and liver segmentation.
202
203    Args:
204        path: Filepath to a folder where the data is downloaded for further processing.
205        batch_size: The batch size for training.
206        patch_shape: The patch shape to use for training.
207        split: The choice of data split.
208        rater: The choice of rater providing the annotations.
209        resize_inputs: Whether to resize inputs to the desired patch shape.
210        download: Whether to download the data if it is not present.
211        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
212
213    Returns:
214        The DataLoader.
215    """
216    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
217    dataset = get_curvas_dataset(path, patch_shape, split, rater, resize_inputs, download, **ds_kwargs)
218    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://zenodo.org/records/12687192/files/training_set.zip'
CHECKSUM = '1126a2205553ae1d4fe5fbaee7ea732aacc4f5a92b96504ed521c23e5a0e3f89'
URLS = {'val': 'https://zenodo.org/records/13767408/files/validation_set.zip', 'test': 'https://zenodo.org/records/13767408/files/testing_set.zip'}
CHECKSUMS = {'val': '01edfac9a085f06111969821d06c83d164654a6041c2e8ac3b11ed390e7c7028', 'test': '6a70aa241a14184778e25d11cae58b39cbc1d8ca204fea19ce0a34bb5b13f7b5'}
H5_DIRS = {'train': 'data', 'val': 'data_val', 'test': 'data_test'}
def get_curvas_data( path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'] = 'train', download: bool = False) -> str:
 78def get_curvas_data(
 79    path: Union[os.PathLike, str], split: Literal["train", "val", "test"] = "train", download: bool = False
 80) -> str:
 81    """Download the CURVAS dataset.
 82
 83    NOTE: The test split is about 21.6 GB.
 84
 85    Args:
 86        path: Filepath to a folder where the data is downloaded for further processing.
 87        split: The choice of data split.
 88        download: Whether to download the data if it is not present.
 89
 90    Returns:
 91        Filepath where the data is downloaded.
 92    """
 93    if split not in H5_DIRS:
 94        raise ValueError(f"'{split}' is not a valid split. Choose one of {list(H5_DIRS)}.")
 95
 96    data_dir = os.path.join(path, H5_DIRS[split])
 97    if os.path.exists(data_dir):
 98        return data_dir
 99
100    os.makedirs(path, exist_ok=True)
101
102    if split == "train":
103        zip_path = os.path.join(path, "training_set.zip")
104        util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
105
106        # HACK: The zip file is broken. We fix it using the following script.
107        fixed_zip_path = os.path.join(path, "training_set_fixed.zip")
108        subprocess.run(["zip", "-FF", zip_path, "--out", fixed_zip_path])
109        subprocess.run(["unzip", fixed_zip_path, "-d", path])
110
111        _preprocess_data(os.path.join(path, "training_set"), data_dir)
112
113        # Remove the zip files as we don't need them anymore.
114        os.remove(zip_path)
115        os.remove(fixed_zip_path)
116    else:
117        zip_path = os.path.join(path, os.path.basename(URLS[split]))
118        util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split])
119        util.unzip(zip_path=zip_path, dst=path, remove=False)
120
121        _preprocess_data(os.path.join(path, os.path.splitext(os.path.basename(URLS[split]))[0]), data_dir)
122
123        os.remove(zip_path)
124
125    return data_dir

Download the CURVAS dataset.

NOTE: The test split is about 21.6 GB.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split.
  • download: Whether to download the data if it is not present.
Returns:

Filepath where the data is downloaded.

def get_curvas_paths( path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False) -> List[str]:
128def get_curvas_paths(
129    path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False
130) -> List[str]:
131    """Get paths to the CURVAS data.
132
133    Args:
134        path: Filepath to a folder where the data is downloaded for further processing.
135        split: The choice of data split.
136        download: Whether to download the data if it is not present.
137
138    Returns:
139        List of filepaths for the volumetric data.
140    """
141    data_dir = get_curvas_data(path, split, download)
142    volume_paths = natsorted(glob(os.path.join(data_dir, "*.h5")))
143
144    if split == "train":
145        volume_paths = volume_paths[:10]
146
147    return volume_paths

Get paths to the CURVAS data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split.
  • download: Whether to download the data if it is not present.
Returns:

List of filepaths for the volumetric data.

def get_curvas_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], split: Literal['train', 'val', 'test'], rater: Literal['1', '2', '3'] = '1', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
150def get_curvas_dataset(
151    path: Union[os.PathLike, str],
152    patch_shape: Tuple[int, ...],
153    split: Literal['train', 'val', 'test'],
154    rater: Literal["1", "2", "3"] = "1",
155    resize_inputs: bool = False,
156    download: bool = False,
157    **kwargs
158) -> Dataset:
159    """Get the CURVAS dataset for pancreas, kidney and liver segmentation.
160
161    Args:
162        path: Filepath to a folder where the data is downloaded for further processing.
163        patch_shape: The patch shape to use for training.
164        split: The choice of data split.
165        rater: The choice of rater providing the annotations.
166        resize_inputs: Whether to resize inputs to the desired patch shape.
167        download: Whether to download the data if it is not present.
168        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
169
170    Returns:
171        The segmentation dataset.
172    """
173    volume_paths = get_curvas_paths(path, split, download)
174
175    if resize_inputs:
176        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
177        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
178            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
179        )
180
181    return torch_em.default_segmentation_dataset(
182        raw_paths=volume_paths,
183        raw_key="raw",
184        label_paths=volume_paths,
185        label_key=f"labels/rater_{rater}",
186        patch_shape=patch_shape,
187        is_seg_dataset=True,
188        **kwargs,
189    )

Get the CURVAS dataset for pancreas, kidney and liver segmentation.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • split: The choice of data split.
  • rater: The choice of rater providing the annotations.
  • 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_curvas_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], split: Literal['train', 'val', 'test'], rater: Literal['1', '2', '3'] = '1', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
192def get_curvas_loader(
193    path: Union[os.PathLike, str],
194    batch_size: int,
195    patch_shape: Tuple[int, ...],
196    split: Literal['train', 'val', 'test'],
197    rater: Literal["1", "2", "3"] = "1",
198    resize_inputs: bool = False,
199    download: bool = False,
200    **kwargs
201) -> DataLoader:
202    """Get the CURVAS dataloader for pancreas, kidney and liver segmentation.
203
204    Args:
205        path: Filepath to a folder where the data is downloaded for further processing.
206        batch_size: The batch size for training.
207        patch_shape: The patch shape to use for training.
208        split: The choice of data split.
209        rater: The choice of rater providing the annotations.
210        resize_inputs: Whether to resize inputs to the desired patch shape.
211        download: Whether to download the data if it is not present.
212        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
213
214    Returns:
215        The DataLoader.
216    """
217    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
218    dataset = get_curvas_dataset(path, patch_shape, split, rater, resize_inputs, download, **ds_kwargs)
219    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the CURVAS dataloader for pancreas, kidney and liver 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.
  • split: The choice of data split.
  • rater: The choice of rater providing the annotations.
  • 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.