torch_em.data.datasets.medical.vs_seg
The Vestibular-Schwannoma-SEG dataset contains annotations for the vestibular schwannoma tumor and the cochlea in MRI.
The dataset consists of contrast-enhanced T1 and high-resolution T2 MRI series of patients undergoing Gamma Knife stereotactic radiosurgery, with expert contours for the tumor (label id 1) and the cochlea (label id 2). A study may have a contour set for either or both of its T1 and T2 series; each RTSTRUCT file is paired with the exact series it references, so both are used where present.
NOTE: The tumor and cochlea ROIs are named inconsistently across the collection (e.g. 'AN', 'TV',
'Rt AN', 'tumour' for the tumor; 'Cochlea', 'cochlea', 'Cochlea_c' for the cochlea), alongside many
unrelated ROIs from serial follow-up measurements (e.g. 'Vol2016', 'Vol 2y') that are not contours of
either structure, so _roi_label maps ROI names to ROI_LABELS by a case-insensitive name match
rather than the exact ROI name.
NOTE: This requires the pydicom python package.
The dataset is located at https://doi.org/10.7937/TCIA.9YTJ-5Q73 and is distributed under the CC BY 4.0 license. This dataset is from the publication https://doi.org/10.1038/s41597-021-01064-w. Please cite it if you use this dataset in your research.
1"""The Vestibular-Schwannoma-SEG dataset contains annotations for the vestibular schwannoma tumor and 2the cochlea in MRI. 3 4The dataset consists of contrast-enhanced T1 and high-resolution T2 MRI series of patients undergoing 5Gamma Knife stereotactic radiosurgery, with expert contours for the tumor (label id 1) and the cochlea 6(label id 2). A study may have a contour set for either or both of its T1 and T2 series; each RTSTRUCT 7file is paired with the exact series it references, so both are used where present. 8 9NOTE: The tumor and cochlea ROIs are named inconsistently across the collection (e.g. 'AN', 'TV', 10'Rt AN', 'tumour' for the tumor; 'Cochlea', 'cochlea', 'Cochlea_c' for the cochlea), alongside many 11unrelated ROIs from serial follow-up measurements (e.g. 'Vol2016', 'Vol 2y') that are not contours of 12either structure, so `_roi_label` maps ROI names to `ROI_LABELS` by a case-insensitive name match 13rather than the exact ROI name. 14 15NOTE: This requires the pydicom python package. 16 17The dataset is located at https://doi.org/10.7937/TCIA.9YTJ-5Q73 and is distributed under the 18CC BY 4.0 license. 19This dataset is from the publication https://doi.org/10.1038/s41597-021-01064-w. 20Please cite it if you use this dataset in your research. 21""" 22 23import os 24import json 25from glob import glob 26from tqdm import tqdm 27from natsort import natsorted 28from typing import Union, Tuple, List 29 30from torch.utils.data import Dataset, DataLoader 31 32import torch_em 33 34from .. import util 35 36 37COLLECTION = "Vestibular-Schwannoma-SEG" 38 39ROI_LABELS = {"tumor": 1, "cochlea": 2} 40"""Mapping from the anatomical structure to its label id.""" 41 42_TUMOR_NAMES = {"an", "rt an", "lt an", "tv", "tvt1", "tumor", "tumour"} 43 44 45def _roi_label(roi_number, roi_name): 46 """Map a ROI to `ROI_LABELS` by a case-insensitive name match, since the collection uses several 47 different names for the same structure and also has unrelated follow-up measurement ROIs. 48 """ 49 name = roi_name.strip().lower() 50 if name in _TUMOR_NAMES: 51 return ROI_LABELS["tumor"] 52 if name.startswith("cochlea"): 53 return ROI_LABELS["cochlea"] 54 return None 55 56 57def _get_series_metadata(path, download): 58 """Get the metadata of all series in the collection from the NBIA REST API.""" 59 import requests 60 61 metadata_path = os.path.join(path, "vs_seg_series.json") 62 if not os.path.exists(metadata_path): 63 if not download: 64 raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.") 65 response = requests.get(f"{util.NBIA_API_URL}getSeries", params={"Collection": COLLECTION}) 66 response.raise_for_status() 67 with open(metadata_path, "w") as f: 68 json.dump(response.json(), f, indent=2) 69 70 with open(metadata_path, "r") as f: 71 return json.load(f) 72 73 74def _referenced_series_uid(rtstruct): 75 referenced_study = rtstruct.ReferencedFrameOfReferenceSequence[0].RTReferencedStudySequence[0] 76 return str(referenced_study.RTReferencedSeriesSequence[0].SeriesInstanceUID) 77 78 79def _preprocess_vs_seg(dicom_dir, series_metadata, preprocessed_dir): 80 import h5py 81 import pydicom 82 83 rtstruct_series = [series for series in series_metadata if series.get("Modality") == "RTSTRUCT"] 84 85 os.makedirs(preprocessed_dir, exist_ok=True) 86 for series in tqdm(rtstruct_series, desc="Preprocess Vestibular-Schwannoma-SEG"): 87 rtstruct_dir = os.path.join(dicom_dir, series["SeriesInstanceUID"]) 88 rtstruct_paths = glob(os.path.join(rtstruct_dir, "*.dcm")) 89 if not rtstruct_paths: 90 continue 91 92 rtstruct = pydicom.dcmread(rtstruct_paths[0], stop_before_pixels=True) 93 mr_uid = _referenced_series_uid(rtstruct) 94 mr_dir = os.path.join(dicom_dir, mr_uid) 95 if not glob(os.path.join(mr_dir, "*.dcm")): 96 continue 97 98 out_path = os.path.join(preprocessed_dir, f"{mr_uid}.h5") 99 if os.path.exists(out_path): 100 continue 101 102 volume, geometry = util.load_dicom_series(mr_dir) 103 labels = util.rasterize_rtstruct(rtstruct_paths[0], geometry, volume.shape, roi_labels=_roi_label) 104 if labels.max() == 0: # A few RTSTRUCT files carry only the auxiliary '*Skull' contour. 105 continue 106 107 with h5py.File(out_path, "w") as f: 108 f.create_dataset("raw", data=volume, compression="gzip") 109 f.create_dataset("labels", data=labels, compression="gzip") 110 111 112def get_vs_seg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 113 """Download the Vestibular-Schwannoma-SEG dataset. 114 115 Args: 116 path: Filepath to a folder where the data is downloaded for further processing. 117 download: Whether to download the data if it is not present. 118 119 Returns: 120 Filepath where the preprocessed data is stored. 121 """ 122 # NOTE: The preprocessing below skips volumes that were converted already, so an interrupted run resumes. 123 preprocessed_dir = os.path.join(path, "preprocessed") 124 125 os.makedirs(path, exist_ok=True) 126 series_metadata = _get_series_metadata(path, download) 127 128 rtstruct_uids = [series["SeriesInstanceUID"] for series in series_metadata if series.get("Modality") == "RTSTRUCT"] 129 130 dicom_dir = os.path.join(path, "dicom") 131 if download: # The RTSTRUCT series are downloaded first, so the MR series they reference can be found. 132 util.download_tcia_series(rtstruct_uids, dst=dicom_dir, csv_filename=os.path.join(path, "vs_seg_rtstruct")) 133 134 import pydicom 135 mr_uids = set() 136 for uid in rtstruct_uids: 137 rtstruct_paths = glob(os.path.join(dicom_dir, uid, "*.dcm")) 138 if rtstruct_paths: 139 rtstruct = pydicom.dcmread(rtstruct_paths[0], stop_before_pixels=True) 140 mr_uids.add(_referenced_series_uid(rtstruct)) 141 util.download_tcia_series(sorted(mr_uids), dst=dicom_dir, csv_filename=os.path.join(path, "vs_seg_mr")) 142 elif not glob(os.path.join(dicom_dir, "*", "*.dcm")): 143 raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.") 144 145 _preprocess_vs_seg(dicom_dir, series_metadata, preprocessed_dir) 146 return preprocessed_dir 147 148 149def get_vs_seg_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]: 150 """Get paths to the Vestibular-Schwannoma-SEG data. 151 152 Args: 153 path: Filepath to a folder where the data is downloaded for further processing. 154 download: Whether to download the data if it is not present. 155 156 Returns: 157 List of filepaths for the stored data. 158 """ 159 preprocessed_dir = get_vs_seg_data(path, download) 160 volume_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.h5"))) 161 assert len(volume_paths) > 0, f"Could not find any preprocessed volumes in '{preprocessed_dir}'." 162 return volume_paths 163 164 165def get_vs_seg_dataset( 166 path: Union[os.PathLike, str], 167 patch_shape: Tuple[int, ...], 168 resize_inputs: bool = False, 169 download: bool = False, 170 **kwargs 171) -> Dataset: 172 """Get the Vestibular-Schwannoma-SEG dataset for tumor and cochlea segmentation. 173 174 Args: 175 path: Filepath to a folder where the data is downloaded for further processing. 176 patch_shape: The patch shape to use for training. 177 resize_inputs: Whether to resize inputs to the desired patch shape. 178 download: Whether to download the data if it is not present. 179 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 180 181 Returns: 182 The segmentation dataset. 183 """ 184 volume_paths = get_vs_seg_paths(path, download) 185 186 if resize_inputs: 187 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 188 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 189 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 190 ) 191 192 return torch_em.default_segmentation_dataset( 193 raw_paths=volume_paths, 194 raw_key="raw", 195 label_paths=volume_paths, 196 label_key="labels", 197 patch_shape=patch_shape, 198 is_seg_dataset=True, 199 **kwargs 200 ) 201 202 203def get_vs_seg_loader( 204 path: Union[os.PathLike, str], 205 batch_size: int, 206 patch_shape: Tuple[int, ...], 207 resize_inputs: bool = False, 208 download: bool = False, 209 **kwargs 210) -> DataLoader: 211 """Get the Vestibular-Schwannoma-SEG dataloader for tumor and cochlea segmentation. 212 213 Args: 214 path: Filepath to a folder where the data is downloaded for further processing. 215 batch_size: The batch size for training. 216 patch_shape: The patch shape to use for training. 217 resize_inputs: Whether to resize inputs to the desired patch shape. 218 download: Whether to download the data if it is not present. 219 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 220 221 Returns: 222 The DataLoader. 223 """ 224 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 225 dataset = get_vs_seg_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 226 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Mapping from the anatomical structure to its label id.
113def get_vs_seg_data(path: Union[os.PathLike, str], download: bool = False) -> str: 114 """Download the Vestibular-Schwannoma-SEG dataset. 115 116 Args: 117 path: Filepath to a folder where the data is downloaded for further processing. 118 download: Whether to download the data if it is not present. 119 120 Returns: 121 Filepath where the preprocessed data is stored. 122 """ 123 # NOTE: The preprocessing below skips volumes that were converted already, so an interrupted run resumes. 124 preprocessed_dir = os.path.join(path, "preprocessed") 125 126 os.makedirs(path, exist_ok=True) 127 series_metadata = _get_series_metadata(path, download) 128 129 rtstruct_uids = [series["SeriesInstanceUID"] for series in series_metadata if series.get("Modality") == "RTSTRUCT"] 130 131 dicom_dir = os.path.join(path, "dicom") 132 if download: # The RTSTRUCT series are downloaded first, so the MR series they reference can be found. 133 util.download_tcia_series(rtstruct_uids, dst=dicom_dir, csv_filename=os.path.join(path, "vs_seg_rtstruct")) 134 135 import pydicom 136 mr_uids = set() 137 for uid in rtstruct_uids: 138 rtstruct_paths = glob(os.path.join(dicom_dir, uid, "*.dcm")) 139 if rtstruct_paths: 140 rtstruct = pydicom.dcmread(rtstruct_paths[0], stop_before_pixels=True) 141 mr_uids.add(_referenced_series_uid(rtstruct)) 142 util.download_tcia_series(sorted(mr_uids), dst=dicom_dir, csv_filename=os.path.join(path, "vs_seg_mr")) 143 elif not glob(os.path.join(dicom_dir, "*", "*.dcm")): 144 raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.") 145 146 _preprocess_vs_seg(dicom_dir, series_metadata, preprocessed_dir) 147 return preprocessed_dir
Download the Vestibular-Schwannoma-SEG 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 preprocessed data is stored.
150def get_vs_seg_paths(path: Union[os.PathLike, str], download: bool = False) -> List[str]: 151 """Get paths to the Vestibular-Schwannoma-SEG data. 152 153 Args: 154 path: Filepath to a folder where the data is downloaded for further processing. 155 download: Whether to download the data if it is not present. 156 157 Returns: 158 List of filepaths for the stored data. 159 """ 160 preprocessed_dir = get_vs_seg_data(path, download) 161 volume_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.h5"))) 162 assert len(volume_paths) > 0, f"Could not find any preprocessed volumes in '{preprocessed_dir}'." 163 return volume_paths
Get paths to the Vestibular-Schwannoma-SEG 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.
166def get_vs_seg_dataset( 167 path: Union[os.PathLike, str], 168 patch_shape: Tuple[int, ...], 169 resize_inputs: bool = False, 170 download: bool = False, 171 **kwargs 172) -> Dataset: 173 """Get the Vestibular-Schwannoma-SEG dataset for tumor and cochlea segmentation. 174 175 Args: 176 path: Filepath to a folder where the data is downloaded for further processing. 177 patch_shape: The patch shape to use for training. 178 resize_inputs: Whether to resize inputs to the desired patch shape. 179 download: Whether to download the data if it is not present. 180 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 181 182 Returns: 183 The segmentation dataset. 184 """ 185 volume_paths = get_vs_seg_paths(path, download) 186 187 if resize_inputs: 188 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 189 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 190 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 191 ) 192 193 return torch_em.default_segmentation_dataset( 194 raw_paths=volume_paths, 195 raw_key="raw", 196 label_paths=volume_paths, 197 label_key="labels", 198 patch_shape=patch_shape, 199 is_seg_dataset=True, 200 **kwargs 201 )
Get the Vestibular-Schwannoma-SEG dataset for tumor and cochlea 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.
204def get_vs_seg_loader( 205 path: Union[os.PathLike, str], 206 batch_size: int, 207 patch_shape: Tuple[int, ...], 208 resize_inputs: bool = False, 209 download: bool = False, 210 **kwargs 211) -> DataLoader: 212 """Get the Vestibular-Schwannoma-SEG dataloader for tumor and cochlea segmentation. 213 214 Args: 215 path: Filepath to a folder where the data is downloaded for further processing. 216 batch_size: The batch size for training. 217 patch_shape: The patch shape to use for training. 218 resize_inputs: Whether to resize inputs to the desired patch shape. 219 download: Whether to download the data if it is not present. 220 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 221 222 Returns: 223 The DataLoader. 224 """ 225 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 226 dataset = get_vs_seg_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 227 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the Vestibular-Schwannoma-SEG dataloader for tumor and cochlea 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.