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