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)
The source collections that require a TCIA Restricted License Agreement and are skipped.
The sentinel that marks a voxel outside the sparsely reviewed slices.
Mapping from the body region name to its label id in body-regions.nii.gz.
Mapping from the body part name to its label id in body-parts.nii.gz.
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.
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.
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' (seeBODY_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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.