torch_em.data.datasets.medical.paxray_pp
The PAXRay++ dataset contains annotations for fine-grained thoracic anatomy segmentation in 2D chest radiographs.
The dataset consists of 14,753 frontal and lateral projections (512x512, generated by projecting
pseudo-labeled thorax CT scans onto a 2D plane) from 7,377 patients, each ('get_paxray_pp_data merges only the 51 'direct' classes
into a single semantic label image per case, where the label id of a structure is its 1-based rank among
the 'direct' channels in the order given by paxray_label_dictionary.json; the 'derived' classes are
redundant with (and would overlap) the 'direct' ones and are not used by this loader.
The data is located at https://doi.org/10.5281/zenodo.20688438, released under a CC-BY-NC-4.0 license.
This dataset is from the publication https://doi.org/10.48550/arXiv.2306.03934. Please cite it if you use this dataset for your research.
1"""The PAXRay++ dataset contains annotations for fine-grained thoracic anatomy segmentation in 22D chest radiographs. 3 4The dataset consists of 14,753 frontal and lateral projections (512x512, generated by projecting 5pseudo-labeled thorax CT scans onto a 2D plane) from 7,377 patients, each ('<case>_frontal.png' or 6'<case>_lateral.png') paired with a boolean multi-channel label volume of shape (159, 512, 512) 7('labels/<case>_<view>.npy'), one channel per anatomical class. Of the 159 classes, 51 are 'direct' 8(directly annotated structures) and 108 are 'derived' (hierarchical unions of the direct structures, 9e.g. 'spine' is the union of 'cervical spine', 'thoracic spine' and 'lumbar spine', which are 10themselves unions of individual vertebrae). `get_paxray_pp_data` merges only the 51 'direct' classes 11into a single semantic label image per case, where the label id of a structure is its 1-based rank among 12the 'direct' channels in the order given by `paxray_label_dictionary.json`; the 'derived' classes are 13redundant with (and would overlap) the 'direct' ones and are not used by this loader. 14 15The data is located at https://doi.org/10.5281/zenodo.20688438, released under a CC-BY-NC-4.0 license. 16 17This dataset is from the publication https://doi.org/10.48550/arXiv.2306.03934. 18Please cite it if you use this dataset for your research. 19""" 20 21import os 22import json 23from glob import glob 24from natsort import natsorted 25from typing import Union, Tuple, List 26 27import numpy as np 28from tqdm import tqdm 29from concurrent import futures 30 31from torch.utils.data import Dataset, DataLoader 32 33import torch_em 34 35from .. import util 36 37 38URLS = { 39 "images": "https://zenodo.org/records/20688438/files/paxray_images.tar.gz", 40 "labels": "https://zenodo.org/records/20688438/files/paxray_labels.zip", 41 "label_dictionary": "https://zenodo.org/records/20688438/files/paxray_label_dictionary.json", 42 "label_types": "https://zenodo.org/records/20688438/files/paxray_label_types.json", 43} 44 45CHECKSUMS = { 46 "images": "e91a9d4d27b5bbfbd24046c897b7facf5391ab934936559fd75b1c92e988b99f", 47 "labels": "9dc4623db0be9793be781699af0605f26a764ab5088efffbb9395adc0e2082c1", 48 "label_dictionary": "f5bf6f48117f18e9cfbc2f844ad7369d043926e32fc4eeae1638b36b05da39a1", 49 "label_types": "465282f18b246a9b8631a145cf7a8d4ccdc4841acea2cc7eb3eedaa4ce633768", 50} 51 52 53def _class_names(path): 54 with open(os.path.join(path, "paxray_label_dictionary.json")) as f: 55 names_by_id = json.load(f) 56 with open(os.path.join(path, "paxray_label_types.json")) as f: 57 types_by_id = json.load(f) 58 59 direct_ids = sorted((int(i) for i in names_by_id if types_by_id[i] == "direct")) 60 return [names_by_id[str(i)] for i in direct_ids], direct_ids 61 62 63def _merge_labels(npy_path, out_path, direct_ids): 64 if os.path.exists(out_path): 65 return out_path 66 67 import tifffile 68 69 channels = np.load(npy_path) 70 label = np.zeros(channels.shape[1:], dtype="uint8") 71 for class_id, channel_id in enumerate(direct_ids, start=1): 72 label[channels[channel_id]] = class_id 73 74 tmp_path = f"{out_path}.incomplete.tif" 75 tifffile.imwrite(tmp_path, label) 76 os.replace(tmp_path, out_path) 77 return out_path 78 79 80def get_paxray_pp_data(path: Union[os.PathLike, str], download: bool = False) -> str: 81 """Download the PAXRay++ dataset and merge the per-class label channels into semantic labels. 82 83 Args: 84 path: Filepath to a folder where the data is downloaded for further processing. 85 download: Whether to download the data if it is not present. 86 87 Returns: 88 Filepath where the data is downloaded. 89 """ 90 image_dir = os.path.join(path, "images", "paxray_images_unfiltered") 91 raw_labels_dir = os.path.join(path, "labels", "labels") 92 semantic_labels_dir = os.path.join(path, "semantic_labels") 93 94 os.makedirs(path, exist_ok=True) 95 96 for name in ("label_dictionary", "label_types"): 97 util.download_source( 98 path=os.path.join(path, f"paxray_{name}.json"), url=URLS[name], download=download, 99 checksum=CHECKSUMS[name], 100 ) 101 _, direct_ids = _class_names(path) 102 103 if not os.path.exists(image_dir): 104 image_tar = os.path.join(path, "images.tar.gz") 105 util.download_source(path=image_tar, url=URLS["images"], download=download, checksum=CHECKSUMS["images"]) 106 util.unzip_tarfile(tar_path=image_tar, dst=os.path.join(path, "images")) 107 108 if not os.path.exists(raw_labels_dir): 109 label_zip = os.path.join(path, "labels.zip") 110 util.download_source(path=label_zip, url=URLS["labels"], download=download, checksum=CHECKSUMS["labels"]) 111 util.unzip(zip_path=label_zip, dst=os.path.join(path, "labels")) 112 113 npy_paths = natsorted(glob(os.path.join(raw_labels_dir, "*.npy"))) 114 if npy_paths and not all( 115 os.path.exists(os.path.join(semantic_labels_dir, f"{os.path.basename(p)[:-4]}.tif")) for p in npy_paths 116 ): 117 os.makedirs(semantic_labels_dir, exist_ok=True) 118 n_workers = min(16, os.cpu_count() or 1) 119 with futures.ProcessPoolExecutor(n_workers) as pool: 120 tasks = [ 121 pool.submit( 122 _merge_labels, p, os.path.join(semantic_labels_dir, f"{os.path.basename(p)[:-4]}.tif"), 123 direct_ids, 124 ) for p in npy_paths 125 ] 126 for task in tqdm(futures.as_completed(tasks), total=len(tasks), desc="Merge the PAXRay++ labels"): 127 task.result() 128 129 return path 130 131 132def get_paxray_pp_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 133 """Get paths to the PAXRay++ data. 134 135 Args: 136 path: Filepath to a folder where the data is downloaded for further processing. 137 download: Whether to download the data if it is not present. 138 139 Returns: 140 List of filepaths for the image data. 141 List of filepaths for the label data. 142 """ 143 data_dir = get_paxray_pp_data(path, download) 144 145 label_paths = natsorted(glob(os.path.join(data_dir, "semantic_labels", "*.tif"))) 146 raw_paths = [ 147 os.path.join(data_dir, "images", "paxray_images_unfiltered", f"{os.path.basename(p)[:-4]}.png") 148 for p in label_paths 149 ] 150 151 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 152 assert all(os.path.exists(p) for p in raw_paths) 153 154 return raw_paths, label_paths 155 156 157def get_paxray_pp_dataset( 158 path: Union[os.PathLike, str], 159 patch_shape: Tuple[int, int], 160 resize_inputs: bool = False, 161 download: bool = False, 162 **kwargs 163) -> Dataset: 164 """Get the PAXRay++ dataset for thoracic anatomy segmentation in chest radiographs. 165 166 Args: 167 path: Filepath to a folder where the data is downloaded for further processing. 168 patch_shape: The patch shape to use for training. 169 resize_inputs: Whether to resize the inputs to the patch shape. 170 download: Whether to download the data if it is not present. 171 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 172 173 Returns: 174 The segmentation dataset. 175 """ 176 raw_paths, label_paths = get_paxray_pp_paths(path, download) 177 178 if resize_inputs: 179 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True} 180 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 181 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 182 ) 183 184 return torch_em.default_segmentation_dataset( 185 raw_paths=raw_paths, 186 raw_key=None, 187 label_paths=label_paths, 188 label_key=None, 189 patch_shape=patch_shape, 190 is_seg_dataset=False, 191 **kwargs 192 ) 193 194 195def get_paxray_pp_loader( 196 path: Union[os.PathLike, str], 197 batch_size: int, 198 patch_shape: Tuple[int, int], 199 resize_inputs: bool = False, 200 download: bool = False, 201 **kwargs 202) -> DataLoader: 203 """Get the PAXRay++ dataloader for thoracic anatomy segmentation in chest radiographs. 204 205 Args: 206 path: Filepath to a folder where the data is downloaded for further processing. 207 batch_size: The batch size for training. 208 patch_shape: The patch shape to use for training. 209 resize_inputs: Whether to resize the inputs to the patch shape. 210 download: Whether to download the data if it is not present. 211 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 212 213 Returns: 214 The DataLoader. 215 """ 216 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 217 dataset = get_paxray_pp_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 218 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
81def get_paxray_pp_data(path: Union[os.PathLike, str], download: bool = False) -> str: 82 """Download the PAXRay++ dataset and merge the per-class label channels into semantic labels. 83 84 Args: 85 path: Filepath to a folder where the data is downloaded for further processing. 86 download: Whether to download the data if it is not present. 87 88 Returns: 89 Filepath where the data is downloaded. 90 """ 91 image_dir = os.path.join(path, "images", "paxray_images_unfiltered") 92 raw_labels_dir = os.path.join(path, "labels", "labels") 93 semantic_labels_dir = os.path.join(path, "semantic_labels") 94 95 os.makedirs(path, exist_ok=True) 96 97 for name in ("label_dictionary", "label_types"): 98 util.download_source( 99 path=os.path.join(path, f"paxray_{name}.json"), url=URLS[name], download=download, 100 checksum=CHECKSUMS[name], 101 ) 102 _, direct_ids = _class_names(path) 103 104 if not os.path.exists(image_dir): 105 image_tar = os.path.join(path, "images.tar.gz") 106 util.download_source(path=image_tar, url=URLS["images"], download=download, checksum=CHECKSUMS["images"]) 107 util.unzip_tarfile(tar_path=image_tar, dst=os.path.join(path, "images")) 108 109 if not os.path.exists(raw_labels_dir): 110 label_zip = os.path.join(path, "labels.zip") 111 util.download_source(path=label_zip, url=URLS["labels"], download=download, checksum=CHECKSUMS["labels"]) 112 util.unzip(zip_path=label_zip, dst=os.path.join(path, "labels")) 113 114 npy_paths = natsorted(glob(os.path.join(raw_labels_dir, "*.npy"))) 115 if npy_paths and not all( 116 os.path.exists(os.path.join(semantic_labels_dir, f"{os.path.basename(p)[:-4]}.tif")) for p in npy_paths 117 ): 118 os.makedirs(semantic_labels_dir, exist_ok=True) 119 n_workers = min(16, os.cpu_count() or 1) 120 with futures.ProcessPoolExecutor(n_workers) as pool: 121 tasks = [ 122 pool.submit( 123 _merge_labels, p, os.path.join(semantic_labels_dir, f"{os.path.basename(p)[:-4]}.tif"), 124 direct_ids, 125 ) for p in npy_paths 126 ] 127 for task in tqdm(futures.as_completed(tasks), total=len(tasks), desc="Merge the PAXRay++ labels"): 128 task.result() 129 130 return path
Download the PAXRay++ dataset and merge the per-class label channels into semantic labels.
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 downloaded.
133def get_paxray_pp_paths(path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]: 134 """Get paths to the PAXRay++ data. 135 136 Args: 137 path: Filepath to a folder where the data is downloaded for further processing. 138 download: Whether to download the data if it is not present. 139 140 Returns: 141 List of filepaths for the image data. 142 List of filepaths for the label data. 143 """ 144 data_dir = get_paxray_pp_data(path, download) 145 146 label_paths = natsorted(glob(os.path.join(data_dir, "semantic_labels", "*.tif"))) 147 raw_paths = [ 148 os.path.join(data_dir, "images", "paxray_images_unfiltered", f"{os.path.basename(p)[:-4]}.png") 149 for p in label_paths 150 ] 151 152 assert len(raw_paths) == len(label_paths) and len(raw_paths) > 0 153 assert all(os.path.exists(p) for p in raw_paths) 154 155 return raw_paths, label_paths
Get paths to the PAXRay++ 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.
158def get_paxray_pp_dataset( 159 path: Union[os.PathLike, str], 160 patch_shape: Tuple[int, int], 161 resize_inputs: bool = False, 162 download: bool = False, 163 **kwargs 164) -> Dataset: 165 """Get the PAXRay++ dataset for thoracic anatomy segmentation in chest radiographs. 166 167 Args: 168 path: Filepath to a folder where the data is downloaded for further processing. 169 patch_shape: The patch shape to use for training. 170 resize_inputs: Whether to resize the inputs to the patch shape. 171 download: Whether to download the data if it is not present. 172 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 173 174 Returns: 175 The segmentation dataset. 176 """ 177 raw_paths, label_paths = get_paxray_pp_paths(path, download) 178 179 if resize_inputs: 180 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True} 181 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 182 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 183 ) 184 185 return torch_em.default_segmentation_dataset( 186 raw_paths=raw_paths, 187 raw_key=None, 188 label_paths=label_paths, 189 label_key=None, 190 patch_shape=patch_shape, 191 is_seg_dataset=False, 192 **kwargs 193 )
Get the PAXRay++ dataset for thoracic anatomy segmentation in chest radiographs.
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 the inputs to the 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.
196def get_paxray_pp_loader( 197 path: Union[os.PathLike, str], 198 batch_size: int, 199 patch_shape: Tuple[int, int], 200 resize_inputs: bool = False, 201 download: bool = False, 202 **kwargs 203) -> DataLoader: 204 """Get the PAXRay++ dataloader for thoracic anatomy segmentation in chest radiographs. 205 206 Args: 207 path: Filepath to a folder where the data is downloaded for further processing. 208 batch_size: The batch size for training. 209 patch_shape: The patch shape to use for training. 210 resize_inputs: Whether to resize the inputs to the patch shape. 211 download: Whether to download the data if it is not present. 212 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 213 214 Returns: 215 The DataLoader. 216 """ 217 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 218 dataset = get_paxray_pp_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs) 219 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the PAXRay++ dataloader for thoracic anatomy segmentation in chest radiographs.
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 the inputs to the 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.