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 ('.change2zip', '.z01' to '.z04'): this module downloads all parts, joins them into a single zip with the 'zip' CLI (Info-ZIP) and extracts it.

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)
KAGGLE_DATASET = 'xiaoweixumedicalai/imagecas'
GROUPS = ['1-200', '201-400', '401-600', '601-800', '801-1000']

The five Kaggle archives that together make up the 1000 cases of the dataset.

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

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

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

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

The DataLoader.