torch_em.data.datasets.medical.imagecas
The ImageCAS dataset contains annotations for coronary artery segmentation in cardiac CT angiography (CCTA).
The dataset consists of 1000 3D CCTA scans collected at the Guangdong Provincial People's Hospital between April 2012 and December 2018. The left and right coronary arteries were independently annotated by two radiologists and cross-verified; disagreements were resolved by a third radiologist. The data is distributed as one 'img.nii.gz' / 'label.nii.gz' pair per case (label id 1 is the coronary artery, 0 is background) and hosted on Kaggle as five zip archives of 200 cases each (https://www.kaggle.com/datasets/xiaoweixumedicalai/imagecas), because the official GitHub repository (https://github.com/XiaoweiXu/ImageCAS-A-Large-Scale-Dataset-and-Benchmark-for-Coronary-Artery-Segmentation-based-on-CT) # noqa does not host the data itself and otherwise requires emailing the authors for access.
Each Kaggle archive is itself split into five parts ('
NOTE: This requires a Kaggle account and API credentials (see https://www.kaggle.com/docs/api), as well as the 'zip' CLI (Info-ZIP) to join the split archives.
This dataset is from the publication https://doi.org/10.1016/j.compmedimag.2023.102287. Please cite it if you use this dataset in your research.
1"""The ImageCAS dataset contains annotations for coronary artery segmentation in cardiac CT angiography (CCTA). 2 3The dataset consists of 1000 3D CCTA scans collected at the Guangdong Provincial People's Hospital between 4April 2012 and December 2018. The left and right coronary arteries were independently annotated by two 5radiologists and cross-verified; disagreements were resolved by a third radiologist. The data is distributed 6as one 'img.nii.gz' / 'label.nii.gz' pair per case (label id 1 is the coronary artery, 0 is background) and 7hosted on Kaggle as five zip archives of 200 cases each (https://www.kaggle.com/datasets/xiaoweixumedicalai/imagecas), 8because the official GitHub repository (https://github.com/XiaoweiXu/ImageCAS-A-Large-Scale-Dataset-and-Benchmark-for-Coronary-Artery-Segmentation-based-on-CT) # noqa 9does not host the data itself and otherwise requires emailing the authors for access. 10 11Each Kaggle archive is itself split into five parts ('<group>.change2zip', '<group>.z01' to '<group>.z04'): 12this module downloads all parts, joins them into a single zip with the 'zip' CLI (Info-ZIP) and extracts it. 13 14NOTE: This requires a Kaggle account and API credentials (see https://www.kaggle.com/docs/api), as well as 15the 'zip' CLI (Info-ZIP) to join the split archives. 16 17This dataset is from the publication https://doi.org/10.1016/j.compmedimag.2023.102287. 18Please cite it if you use this dataset in your research. 19""" 20 21import os 22from glob import glob 23from shutil import which 24from subprocess import run 25from natsort import natsorted 26from typing import Union, Tuple, List 27 28from torch.utils.data import Dataset, DataLoader 29 30import torch_em 31 32from .. import util 33 34 35KAGGLE_DATASET = "xiaoweixumedicalai/imagecas" 36 37GROUPS = ["1-200", "201-400", "401-600", "601-800", "801-1000"] 38"""The five Kaggle archives that together make up the 1000 cases of the dataset.""" 39 40 41def _download_kaggle_file(filename: str, dst_dir: str, download: bool) -> str: 42 """Download a single file from the ImageCAS Kaggle dataset. 43 44 Kaggle wraps every single-file download in an outer zip container (even if the file is itself 45 already an archive), which is unpacked here to recover the original file. 46 """ 47 out_path = os.path.join(dst_dir, filename) 48 if os.path.exists(out_path): 49 return out_path 50 if not download: 51 raise RuntimeError(f"Cannot find the data at {out_path}, but download was set to False.") 52 53 try: 54 from kaggle.api.kaggle_api_extended import KaggleApi 55 except ModuleNotFoundError: 56 msg = "Please install the Kaggle API. You can do this using 'pip install kaggle'. " 57 msg += "After you have installed kaggle, you would need an API token. " 58 msg += "Follow the instructions at https://www.kaggle.com/docs/api." 59 raise ModuleNotFoundError(msg) 60 61 os.makedirs(dst_dir, exist_ok=True) 62 api = KaggleApi() 63 api.authenticate() 64 api.dataset_download_file(KAGGLE_DATASET, filename, path=dst_dir) 65 66 wrapper_path = os.path.join(dst_dir, f"{filename}.zip") 67 util.unzip(zip_path=wrapper_path, dst=dst_dir) 68 return out_path 69 70 71def _rename_nifti_files(case_dir: str) -> None: 72 """Rename '<id>.img.nii.gz' / '<id>.label.nii.gz' to '<id>_img.nii.gz' / '<id>_label.nii.gz'. 73 74 'elf.io.open_file' (used by `torch_em.data.SegmentationDataset`) only recognizes '.nii.gz' files that have 75 exactly two suffixes, e.g. '<id>.nii.gz'. The extra '.img' / '.label' suffix in the original file names would 76 otherwise be mistaken for the file extension, so the files are renamed once after extraction. 77 """ 78 for suffix in ("img", "label"): 79 for path in glob(os.path.join(case_dir, f"*.{suffix}.nii.gz")): 80 new_path = path[:-len(f".{suffix}.nii.gz")] + f"_{suffix}.nii.gz" 81 if not os.path.exists(new_path): 82 os.rename(path, new_path) 83 84 85def _merge_and_extract_group(group: str, zip_dir: str, raw_dir: str, download: bool) -> None: 86 """Download, join and extract the split zip archive of one group (200 cases) of the dataset. 87 88 Groups that were already extracted (e.g. by a previous, interrupted run) are skipped. 89 """ 90 if glob(os.path.join(raw_dir, group, "*_img.nii.gz")): 91 return 92 93 parts = [f"{group}.change2zip"] + [f"{group}.z0{i}" for i in range(1, 5)] 94 for part in parts: 95 _download_kaggle_file(part, zip_dir, download) 96 97 base_zip = os.path.join(zip_dir, f"{group}.zip") 98 if not os.path.exists(base_zip): 99 os.rename(os.path.join(zip_dir, f"{group}.change2zip"), base_zip) 100 101 merged_zip = os.path.join(zip_dir, f"{group}.merged.zip") 102 if not os.path.exists(merged_zip): 103 if which("zip") is None: 104 raise RuntimeError( 105 "Need the 'zip' CLI (Info-ZIP) to join the split zip archive of the ImageCAS dataset. " 106 "You can install it via 'conda install -c conda-forge zip'." 107 ) 108 run(["zip", "-s", "0", base_zip, "--out", merged_zip], check=True, cwd=zip_dir) 109 110 util.unzip(zip_path=merged_zip, dst=raw_dir, remove=False) 111 _rename_nifti_files(os.path.join(raw_dir, group)) 112 113 # The split zip parts and the joined zip are removed once the group has been extracted, so that the ~18 GB 114 # per group of intermediate files do not pile up on disk (the extraction itself is not repeated afterwards). 115 for part in parts[1:]: 116 part_path = os.path.join(zip_dir, part) 117 if os.path.exists(part_path): 118 os.remove(part_path) 119 for leftover in (base_zip, merged_zip): 120 if os.path.exists(leftover): 121 os.remove(leftover) 122 123 124def get_imagecas_data(path: Union[os.PathLike, str], download: bool = False) -> str: 125 """Download the ImageCAS dataset. 126 127 Args: 128 path: Filepath to a folder where the data is downloaded for further processing. 129 download: Whether to download the data if it is not present. 130 131 Returns: 132 Filepath where the data is stored. 133 """ 134 raw_dir = os.path.join(path, "data") 135 if len(glob(os.path.join(raw_dir, "**", "*_img.nii.gz"), recursive=True)) >= 1000: 136 return raw_dir 137 138 os.makedirs(raw_dir, exist_ok=True) 139 140 zip_dir = os.path.join(path, "zips") 141 for group in GROUPS: 142 _merge_and_extract_group(group, zip_dir, raw_dir, download) 143 144 return raw_dir 145 146 147def get_imagecas_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 148 """Get paths to the ImageCAS data. 149 150 Args: 151 path: Filepath to a folder where the data is downloaded for further processing. 152 download: Whether to download the data if it is not present. 153 154 Returns: 155 List of filepaths for the image data. 156 List of filepaths for the label data. 157 """ 158 raw_dir = get_imagecas_data(path, download) 159 160 image_paths = natsorted(glob(os.path.join(raw_dir, "**", "*_img.nii.gz"), recursive=True)) 161 label_paths = natsorted(glob(os.path.join(raw_dir, "**", "*_label.nii.gz"), recursive=True)) 162 assert len(image_paths) > 0 and len(image_paths) == len(label_paths), \ 163 f"Could not find a matching number of images and labels in '{raw_dir}'." 164 165 return image_paths, label_paths 166 167 168def get_imagecas_dataset( 169 path: Union[os.PathLike, str], 170 patch_shape: Tuple[int, ...], 171 resize_inputs: bool = False, 172 download: bool = False, 173 **kwargs 174) -> Dataset: 175 """Get the ImageCAS dataset for coronary artery segmentation. 176 177 Args: 178 path: Filepath to a folder where the data is downloaded for further processing. 179 patch_shape: The patch shape to use for training. 180 resize_inputs: Whether to resize inputs to the desired patch shape. 181 download: Whether to download the data if it is not present. 182 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 183 184 Returns: 185 The segmentation dataset. 186 """ 187 image_paths, label_paths = get_imagecas_paths(path, download) 188 189 if resize_inputs: 190 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 191 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 192 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 193 ) 194 195 return torch_em.default_segmentation_dataset( 196 raw_paths=image_paths, 197 raw_key="data", 198 label_paths=label_paths, 199 label_key="data", 200 patch_shape=patch_shape, 201 is_seg_dataset=True, 202 **kwargs 203 ) 204 205 206def get_imagecas_loader( 207 path: Union[os.PathLike, str], 208 batch_size: int, 209 patch_shape: Tuple[int, ...], 210 resize_inputs: bool = False, 211 download: bool = False, 212 **kwargs 213) -> DataLoader: 214 """Get the ImageCAS dataloader for coronary artery segmentation. 215 216 Args: 217 path: Filepath to a folder where the data is downloaded for further processing. 218 batch_size: The batch size for training. 219 patch_shape: The patch shape to use for training. 220 resize_inputs: Whether to resize inputs to the desired patch shape. 221 download: Whether to download the data if it is not present. 222 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 223 224 Returns: 225 The DataLoader. 226 """ 227 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 228 dataset = get_imagecas_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 229 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
The five Kaggle archives that together make up the 1000 cases of the dataset.
125def get_imagecas_data(path: Union[os.PathLike, str], download: bool = False) -> str: 126 """Download the ImageCAS dataset. 127 128 Args: 129 path: Filepath to a folder where the data is downloaded for further processing. 130 download: Whether to download the data if it is not present. 131 132 Returns: 133 Filepath where the data is stored. 134 """ 135 raw_dir = os.path.join(path, "data") 136 if len(glob(os.path.join(raw_dir, "**", "*_img.nii.gz"), recursive=True)) >= 1000: 137 return raw_dir 138 139 os.makedirs(raw_dir, exist_ok=True) 140 141 zip_dir = os.path.join(path, "zips") 142 for group in GROUPS: 143 _merge_and_extract_group(group, zip_dir, raw_dir, download) 144 145 return raw_dir
Download the ImageCAS 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 stored.
148def get_imagecas_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 149 """Get paths to the ImageCAS data. 150 151 Args: 152 path: Filepath to a folder where the data is downloaded for further processing. 153 download: Whether to download the data if it is not present. 154 155 Returns: 156 List of filepaths for the image data. 157 List of filepaths for the label data. 158 """ 159 raw_dir = get_imagecas_data(path, download) 160 161 image_paths = natsorted(glob(os.path.join(raw_dir, "**", "*_img.nii.gz"), recursive=True)) 162 label_paths = natsorted(glob(os.path.join(raw_dir, "**", "*_label.nii.gz"), recursive=True)) 163 assert len(image_paths) > 0 and len(image_paths) == len(label_paths), \ 164 f"Could not find a matching number of images and labels in '{raw_dir}'." 165 166 return image_paths, label_paths
Get paths to the ImageCAS 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.
169def get_imagecas_dataset( 170 path: Union[os.PathLike, str], 171 patch_shape: Tuple[int, ...], 172 resize_inputs: bool = False, 173 download: bool = False, 174 **kwargs 175) -> Dataset: 176 """Get the ImageCAS dataset for coronary artery segmentation. 177 178 Args: 179 path: Filepath to a folder where the data is downloaded for further processing. 180 patch_shape: The patch shape to use for training. 181 resize_inputs: Whether to resize inputs to the desired patch shape. 182 download: Whether to download the data if it is not present. 183 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 184 185 Returns: 186 The segmentation dataset. 187 """ 188 image_paths, label_paths = get_imagecas_paths(path, download) 189 190 if resize_inputs: 191 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 192 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 193 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 194 ) 195 196 return torch_em.default_segmentation_dataset( 197 raw_paths=image_paths, 198 raw_key="data", 199 label_paths=label_paths, 200 label_key="data", 201 patch_shape=patch_shape, 202 is_seg_dataset=True, 203 **kwargs 204 )
Get the ImageCAS dataset for coronary artery 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.
207def get_imagecas_loader( 208 path: Union[os.PathLike, str], 209 batch_size: int, 210 patch_shape: Tuple[int, ...], 211 resize_inputs: bool = False, 212 download: bool = False, 213 **kwargs 214) -> DataLoader: 215 """Get the ImageCAS dataloader for coronary artery segmentation. 216 217 Args: 218 path: Filepath to a folder where the data is downloaded for further processing. 219 batch_size: The batch size for training. 220 patch_shape: The patch shape to use for training. 221 resize_inputs: Whether to resize inputs to the desired patch shape. 222 download: Whether to download the data if it is not present. 223 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 224 225 Returns: 226 The DataLoader. 227 """ 228 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 229 dataset = get_imagecas_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 230 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the ImageCAS dataloader for coronary artery 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.