torch_em.data.datasets.medical.saros

The SAROS dataset contains annotations for 13 body regions and 6 body parts in whole-body CT.

The dataset consists of 900 CT series pooled from 28 TCIA collections, each resampled to 5mm slice thickness and given two label volumes on that same grid: body-regions.nii.gz (see BODY_REGIONS) and body-parts.nii.gz (see BODY_PARTS). Both are sparsely annotated: only every 5th axial slice was reviewed by an annotator, and IGNORE_LABEL marks every other slice.

NOTE: The images are not distributed with the release: only the label volumes and a manifest CSV are, so this module downloads and reconstructs them from their original TCIA series, following the same steps and settings as the release. A raw DICOM conversion (e.g. with dcm2niix) does not share the label's grid, so it is resampled onto it: DICOM patient coordinates are LPS, the label is stored in a RAS+ world frame, and a trilinear resampling with a -1024 HU fill value outside the CT extent completes the match.

NOTE: 6 of the 28 source collections (Head-Neck Cetuximab, ACRIN-HNSCC-FDG-PET-CT, QIN-HEADNECK, TCGA-HNSC, HNSCC, Anti-PD-1_MELANOMA) require signing a TCIA Restricted License Agreement and are not reachable through the public NBIA API, so their 174 cases are skipped; the remaining 726 are openly downloadable.

NOTE: This requires the pydicom, nibabel and scipy python packages.

The dataset is located at https://doi.org/10.25737/sz96-zg60 and is distributed under the TCIA Restricted License / CC BY 4.0 license (per-collection, see the collection's own citation). This dataset is from the publication https://doi.org/10.1038/s41597-024-03337-6. Please cite it if you use this dataset in your research.

  1"""The SAROS dataset contains annotations for 13 body regions and 6 body parts in whole-body CT.
  2
  3The dataset consists of 900 CT series pooled from 28 TCIA collections, each resampled to 5mm slice
  4thickness and given two label volumes on that same grid: `body-regions.nii.gz` (see `BODY_REGIONS`)
  5and `body-parts.nii.gz` (see `BODY_PARTS`). Both are sparsely annotated: only every 5th axial slice
  6was reviewed by an annotator, and `IGNORE_LABEL` marks every other slice.
  7
  8NOTE: The images are not distributed with the release: only the label volumes and a manifest CSV
  9are, so this module downloads and reconstructs them from their original TCIA series, following the
 10same steps and settings as the release. A raw DICOM conversion (e.g. with dcm2niix) does not share
 11the label's grid, so it is resampled onto it: DICOM patient coordinates are LPS, the label is stored
 12in a RAS+ world frame, and a trilinear resampling with a -1024 HU fill value outside the CT extent
 13completes the match.
 14
 15NOTE: 6 of the 28 source collections (Head-Neck Cetuximab, ACRIN-HNSCC-FDG-PET-CT, QIN-HEADNECK,
 16TCGA-HNSC, HNSCC, Anti-PD-1_MELANOMA) require signing a TCIA Restricted License Agreement and are
 17not reachable through the public NBIA API, so their 174 cases are skipped; the remaining 726 are
 18openly downloadable.
 19
 20NOTE: This requires the pydicom, nibabel and scipy python packages.
 21
 22The dataset is located at https://doi.org/10.25737/sz96-zg60 and is distributed under the
 23TCIA Restricted License / CC BY 4.0 license (per-collection, see the collection's own citation).
 24This dataset is from the publication https://doi.org/10.1038/s41597-024-03337-6.
 25Please cite it if you use this dataset in your research.
 26"""
 27
 28import os
 29import csv
 30from glob import glob
 31from tqdm import tqdm
 32from natsort import natsorted
 33from typing import Union, Tuple, List
 34
 35import numpy as np
 36
 37from torch.utils.data import Dataset, DataLoader
 38
 39import torch_em
 40
 41from .adrenal_acc import _load_dicom_volume
 42from .. import util
 43
 44
 45URLS = {
 46    "segs": "https://www.cancerimagingarchive.net/wp-content/uploads/SAROS-Collection-NIfTI-files-v2_03-70-2024.zip",  # noqa
 47    "info": "https://www.cancerimagingarchive.net/wp-content/uploads/Segmentation-Info_09-29-2023.csv",
 48}
 49
 50CHECKSUMS = {
 51    "segs": "b509ff70fa69673b0697dac711a92b0e04476780feadf089927a7c8fcd7037e5",
 52    "info": "dac6df664279965567b79ff816a23d6f851cd7ab23340e81ff581c9b079c0cb1",
 53}
 54
 55RESTRICTED_COLLECTIONS = {
 56    "Head-Neck Cetuximab", "ACRIN-HNSCC-FDG-PET-CT", "QIN-HEADNECK", "TCGA-HNSC", "HNSCC", "Anti-PD-1_MELANOMA",
 57}
 58"""The source collections that require a TCIA Restricted License Agreement and are skipped."""
 59
 60IGNORE_LABEL = 255
 61"""The sentinel that marks a voxel outside the sparsely reviewed slices."""
 62
 63BODY_REGIONS = {
 64    "subcutaneous_tissue": 1, "muscle": 2, "abdominal_cavity": 3, "thoracic_cavity": 4, "bone": 5,
 65    "parotid_glands": 6, "pericardium": 7, "breast_implant": 8, "mediastinum": 9, "brain": 10,
 66    "spinal_cord": 11, "thyroid_glands": 12, "submandibular_glands": 13,
 67}
 68"""Mapping from the body region name to its label id in `body-regions.nii.gz`."""
 69
 70BODY_PARTS = {"torso": 1, "head": 2, "right_leg": 3, "left_leg": 4, "right_arm": 5, "left_arm": 6}
 71"""Mapping from the body part name to its label id in `body-parts.nii.gz`."""
 72
 73
 74def _read_manifest(info_path):
 75    with open(info_path) as f:
 76        return list(csv.DictReader(f))
 77
 78
 79def _resample_to_label(volume, ct_affine, label_shape, label_affine):
 80    """Resample a DICOM-derived volume onto the grid of its label, matching the release's own
 81    reconstruction: DICOM patient coordinates are LPS, converted to the RAS+ frame of the label by
 82    negating x and y, then a trilinear resampling with a -1024 HU fill value outside the CT extent.
 83    """
 84    from scipy.ndimage import affine_transform
 85
 86    lps_to_ras = np.diag([-1.0, -1.0, 1.0, 1.0])
 87    ras_affine = lps_to_ras @ ct_affine
 88    if volume.shape == label_shape and np.allclose(ras_affine, label_affine, atol=1e-3):
 89        return volume.astype("int16")
 90
 91    to_ct_index = np.linalg.inv(ras_affine) @ label_affine
 92    resampled = affine_transform(
 93        volume.astype("float32"), to_ct_index[:3, :3], offset=to_ct_index[:3, 3],
 94        output_shape=label_shape, order=1, mode="constant", cval=-1024.0,
 95    )
 96    return np.round(resampled).astype("int16")
 97
 98
 99def _preprocess_saros(seg_dir, manifest, dicom_dir, preprocessed_dir):
100    import h5py
101    import nibabel as nib
102
103    os.makedirs(preprocessed_dir, exist_ok=True)
104    for row in tqdm(manifest, desc="Preprocess SAROS"):
105        case_id = row["id"]
106        out_path = os.path.join(preprocessed_dir, f"{case_id}.h5")
107        if os.path.exists(out_path):
108            continue
109
110        regions_path = os.path.join(seg_dir, case_id, "body-regions.nii.gz")
111        parts_path = os.path.join(seg_dir, case_id, "body-parts.nii.gz")
112        if not (os.path.exists(regions_path) and os.path.exists(parts_path)):
113            continue
114
115        series_dir = os.path.join(dicom_dir, row["tcia_series_instance_uid"])
116        if not glob(os.path.join(series_dir, "*.dcm")):
117            continue
118
119        regions_image = nib.load(regions_path)
120        regions = np.asarray(regions_image.dataobj)
121        parts = np.asarray(nib.load(parts_path).dataobj)
122
123        volume, ct_affine = _load_dicom_volume(series_dir)
124        raw = _resample_to_label(volume, ct_affine, regions.shape, regions_image.affine)
125
126        with h5py.File(out_path, "w") as f:
127            f.create_dataset("raw", data=raw, compression="gzip")
128            f.create_dataset("labels/regions", data=regions.astype("uint8"), compression="gzip")
129            f.create_dataset("labels/parts", data=parts.astype("uint8"), compression="gzip")
130
131
132def get_saros_data(path: Union[os.PathLike, str], download: bool = False) -> str:
133    """Download the SAROS dataset.
134
135    The images are reconstructed from TCIA, which is several hundred gigabytes and can take many
136    hours to download depending on the connection to the TCIA servers.
137
138    Args:
139        path: Filepath to a folder where the data is downloaded for further processing.
140        download: Whether to download the data if it is not present.
141
142    Returns:
143        Filepath where the preprocessed data is stored.
144    """
145    # NOTE: The preprocessing below skips volumes that were converted already, so an interrupted run resumes.
146    preprocessed_dir = os.path.join(path, "preprocessed")
147
148    os.makedirs(path, exist_ok=True)
149
150    info_path = os.path.join(path, "info.csv")
151    util.download_source(path=info_path, url=URLS["info"], download=download, checksum=CHECKSUMS["info"])
152
153    seg_dir = os.path.join(path, "segs")
154    if not os.path.exists(seg_dir):
155        zip_path = os.path.join(path, "segs.zip")
156        util.download_source(path=zip_path, url=URLS["segs"], download=download, checksum=CHECKSUMS["segs"])
157        util.unzip(zip_path=zip_path, dst=seg_dir, remove=False)
158
159    manifest = [row for row in _read_manifest(info_path) if row["tcia_collection"] not in RESTRICTED_COLLECTIONS]
160    series_uids = sorted({row["tcia_series_instance_uid"] for row in manifest})
161
162    dicom_dir = os.path.join(path, "dicom")
163    if download:  # Series that were downloaded already are skipped.
164        util.download_tcia_series(series_uids, dst=dicom_dir, csv_filename=os.path.join(path, "saros"))
165    elif not all(glob(os.path.join(dicom_dir, uid, "*.dcm")) for uid in series_uids):
166        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
167
168    # The extracted collection nests the per-case folders one level deeper (case_XXX / body-*.nii.gz).
169    case_dirs = glob(os.path.join(seg_dir, "*", "case_*"))
170    seg_root = os.path.dirname(case_dirs[0]) if case_dirs else seg_dir
171
172    _preprocess_saros(seg_root, manifest, dicom_dir, preprocessed_dir)
173    return preprocessed_dir
174
175
176def get_saros_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
177    """Get paths to the SAROS data.
178
179    Args:
180        path: Filepath to a folder where the data is downloaded for further processing.
181        download: Whether to download the data if it is not present.
182
183    Returns:
184        List of filepaths for the stored data.
185    """
186    preprocessed_dir = get_saros_data(path, download)
187    volume_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.h5")))
188    assert len(volume_paths) > 0, f"Could not find any preprocessed volumes in '{preprocessed_dir}'."
189    return volume_paths
190
191
192def get_saros_dataset(
193    path: Union[os.PathLike, str],
194    patch_shape: Tuple[int, ...],
195    label_type: str = "regions",
196    resize_inputs: bool = False,
197    download: bool = False,
198    **kwargs
199) -> Dataset:
200    """Get the SAROS dataset for body region or body part segmentation.
201
202    Args:
203        path: Filepath to a folder where the data is downloaded for further processing.
204        patch_shape: The patch shape to use for training.
205        label_type: The label volume to use, one of 'regions' (see `BODY_REGIONS`) or 'parts'
206            (see `BODY_PARTS`).
207        resize_inputs: Whether to resize inputs to the desired patch shape.
208        download: Whether to download the data if it is not present.
209        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
210
211    Returns:
212        The segmentation dataset.
213    """
214    assert label_type in ("regions", "parts"), f"'{label_type}' is not a valid label type."
215    volume_paths = get_saros_paths(path, download)
216
217    if resize_inputs:
218        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
219        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
220            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
221        )
222
223    return torch_em.default_segmentation_dataset(
224        raw_paths=volume_paths,
225        raw_key="raw",
226        label_paths=volume_paths,
227        label_key=f"labels/{label_type}",
228        patch_shape=patch_shape,
229        is_seg_dataset=True,
230        **kwargs
231    )
232
233
234def get_saros_loader(
235    path: Union[os.PathLike, str],
236    batch_size: int,
237    patch_shape: Tuple[int, ...],
238    label_type: str = "regions",
239    resize_inputs: bool = False,
240    download: bool = False,
241    **kwargs
242) -> DataLoader:
243    """Get the SAROS dataloader for body region or body part segmentation.
244
245    Args:
246        path: Filepath to a folder where the data is downloaded for further processing.
247        batch_size: The batch size for training.
248        patch_shape: The patch shape to use for training.
249        label_type: The label volume to use, one of 'regions' or 'parts'.
250        resize_inputs: Whether to resize inputs to the desired patch shape.
251        download: Whether to download the data if it is not present.
252        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
253
254    Returns:
255        The DataLoader.
256    """
257    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
258    dataset = get_saros_dataset(path, patch_shape, label_type, resize_inputs, download, **ds_kwargs)
259    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URLS = {'segs': 'https://www.cancerimagingarchive.net/wp-content/uploads/SAROS-Collection-NIfTI-files-v2_03-70-2024.zip', 'info': 'https://www.cancerimagingarchive.net/wp-content/uploads/Segmentation-Info_09-29-2023.csv'}
CHECKSUMS = {'segs': 'b509ff70fa69673b0697dac711a92b0e04476780feadf089927a7c8fcd7037e5', 'info': 'dac6df664279965567b79ff816a23d6f851cd7ab23340e81ff581c9b079c0cb1'}
RESTRICTED_COLLECTIONS = {'Anti-PD-1_MELANOMA', 'TCGA-HNSC', 'HNSCC', 'QIN-HEADNECK', 'ACRIN-HNSCC-FDG-PET-CT', 'Head-Neck Cetuximab'}

The source collections that require a TCIA Restricted License Agreement and are skipped.

IGNORE_LABEL = 255

The sentinel that marks a voxel outside the sparsely reviewed slices.

BODY_REGIONS = {'subcutaneous_tissue': 1, 'muscle': 2, 'abdominal_cavity': 3, 'thoracic_cavity': 4, 'bone': 5, 'parotid_glands': 6, 'pericardium': 7, 'breast_implant': 8, 'mediastinum': 9, 'brain': 10, 'spinal_cord': 11, 'thyroid_glands': 12, 'submandibular_glands': 13}

Mapping from the body region name to its label id in body-regions.nii.gz.

BODY_PARTS = {'torso': 1, 'head': 2, 'right_leg': 3, 'left_leg': 4, 'right_arm': 5, 'left_arm': 6}

Mapping from the body part name to its label id in body-parts.nii.gz.

def get_saros_data(path: Union[os.PathLike, str], download: bool = False) -> str:
133def get_saros_data(path: Union[os.PathLike, str], download: bool = False) -> str:
134    """Download the SAROS dataset.
135
136    The images are reconstructed from TCIA, which is several hundred gigabytes and can take many
137    hours to download depending on the connection to the TCIA servers.
138
139    Args:
140        path: Filepath to a folder where the data is downloaded for further processing.
141        download: Whether to download the data if it is not present.
142
143    Returns:
144        Filepath where the preprocessed data is stored.
145    """
146    # NOTE: The preprocessing below skips volumes that were converted already, so an interrupted run resumes.
147    preprocessed_dir = os.path.join(path, "preprocessed")
148
149    os.makedirs(path, exist_ok=True)
150
151    info_path = os.path.join(path, "info.csv")
152    util.download_source(path=info_path, url=URLS["info"], download=download, checksum=CHECKSUMS["info"])
153
154    seg_dir = os.path.join(path, "segs")
155    if not os.path.exists(seg_dir):
156        zip_path = os.path.join(path, "segs.zip")
157        util.download_source(path=zip_path, url=URLS["segs"], download=download, checksum=CHECKSUMS["segs"])
158        util.unzip(zip_path=zip_path, dst=seg_dir, remove=False)
159
160    manifest = [row for row in _read_manifest(info_path) if row["tcia_collection"] not in RESTRICTED_COLLECTIONS]
161    series_uids = sorted({row["tcia_series_instance_uid"] for row in manifest})
162
163    dicom_dir = os.path.join(path, "dicom")
164    if download:  # Series that were downloaded already are skipped.
165        util.download_tcia_series(series_uids, dst=dicom_dir, csv_filename=os.path.join(path, "saros"))
166    elif not all(glob(os.path.join(dicom_dir, uid, "*.dcm")) for uid in series_uids):
167        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
168
169    # The extracted collection nests the per-case folders one level deeper (case_XXX / body-*.nii.gz).
170    case_dirs = glob(os.path.join(seg_dir, "*", "case_*"))
171    seg_root = os.path.dirname(case_dirs[0]) if case_dirs else seg_dir
172
173    _preprocess_saros(seg_root, manifest, dicom_dir, preprocessed_dir)
174    return preprocessed_dir

Download the SAROS dataset.

The images are reconstructed from TCIA, which is several hundred gigabytes and can take many hours to download depending on the connection to the TCIA servers.

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 preprocessed data is stored.

def get_saros_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
177def get_saros_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]:
178    """Get paths to the SAROS data.
179
180    Args:
181        path: Filepath to a folder where the data is downloaded for further processing.
182        download: Whether to download the data if it is not present.
183
184    Returns:
185        List of filepaths for the stored data.
186    """
187    preprocessed_dir = get_saros_data(path, download)
188    volume_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.h5")))
189    assert len(volume_paths) > 0, f"Could not find any preprocessed volumes in '{preprocessed_dir}'."
190    return volume_paths

Get paths to the SAROS 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 stored data.

def get_saros_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], label_type: str = 'regions', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
193def get_saros_dataset(
194    path: Union[os.PathLike, str],
195    patch_shape: Tuple[int, ...],
196    label_type: str = "regions",
197    resize_inputs: bool = False,
198    download: bool = False,
199    **kwargs
200) -> Dataset:
201    """Get the SAROS dataset for body region or body part segmentation.
202
203    Args:
204        path: Filepath to a folder where the data is downloaded for further processing.
205        patch_shape: The patch shape to use for training.
206        label_type: The label volume to use, one of 'regions' (see `BODY_REGIONS`) or 'parts'
207            (see `BODY_PARTS`).
208        resize_inputs: Whether to resize inputs to the desired patch shape.
209        download: Whether to download the data if it is not present.
210        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
211
212    Returns:
213        The segmentation dataset.
214    """
215    assert label_type in ("regions", "parts"), f"'{label_type}' is not a valid label type."
216    volume_paths = get_saros_paths(path, download)
217
218    if resize_inputs:
219        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
220        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
221            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
222        )
223
224    return torch_em.default_segmentation_dataset(
225        raw_paths=volume_paths,
226        raw_key="raw",
227        label_paths=volume_paths,
228        label_key=f"labels/{label_type}",
229        patch_shape=patch_shape,
230        is_seg_dataset=True,
231        **kwargs
232    )

Get the SAROS dataset for body region or body part segmentation.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • label_type: The label volume to use, one of 'regions' (see BODY_REGIONS) or 'parts' (see BODY_PARTS).
  • 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_saros_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], label_type: str = 'regions', resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
235def get_saros_loader(
236    path: Union[os.PathLike, str],
237    batch_size: int,
238    patch_shape: Tuple[int, ...],
239    label_type: str = "regions",
240    resize_inputs: bool = False,
241    download: bool = False,
242    **kwargs
243) -> DataLoader:
244    """Get the SAROS dataloader for body region or body part segmentation.
245
246    Args:
247        path: Filepath to a folder where the data is downloaded for further processing.
248        batch_size: The batch size for training.
249        patch_shape: The patch shape to use for training.
250        label_type: The label volume to use, one of 'regions' or 'parts'.
251        resize_inputs: Whether to resize inputs to the desired patch shape.
252        download: Whether to download the data if it is not present.
253        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
254
255    Returns:
256        The DataLoader.
257    """
258    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
259    dataset = get_saros_dataset(path, patch_shape, label_type, resize_inputs, download, **ds_kwargs)
260    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the SAROS dataloader for body region or body part 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.
  • label_type: The label volume to use, one of 'regions' or 'parts'.
  • 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.