torch_em.data.datasets.medical.chestx_det
The ChestX-Det dataset contains annotations for segmentation of 13 categories of thoracic abnormalities or diseases in chest x-ray images.
The dataset consists of 3578 images from NIH ChestX-14, annotated by three board-certified
radiologists with polygon contours for the 13 categories (see CHESTX_DET_LABELS). The dataset
is located at https://github.com/Deepwise-AILab/ChestX-Det-Dataset and is distributed under the
Apache 2.0 license.
This dataset is from the publication https://doi.org/10.48550/arXiv.2004.10871. Please cite it if you use this dataset for your research.
1"""The ChestX-Det dataset contains annotations for segmentation of 13 categories of thoracic 2abnormalities or diseases in chest x-ray images. 3 4The dataset consists of 3578 images from NIH ChestX-14, annotated by three board-certified 5radiologists with polygon contours for the 13 categories (see `CHESTX_DET_LABELS`). The dataset 6is located at https://github.com/Deepwise-AILab/ChestX-Det-Dataset and is distributed under the 7Apache 2.0 license. 8 9This dataset is from the publication https://doi.org/10.48550/arXiv.2004.10871. 10Please cite it if you use this dataset for your research. 11""" 12 13import os 14import json 15from glob import glob 16from tqdm import tqdm 17from natsort import natsorted 18from typing import Union, Tuple, List, Literal 19 20import numpy as np 21from skimage.draw import polygon 22import imageio.v3 as imageio 23 24from torch.utils.data import Dataset, DataLoader 25 26import torch_em 27 28from .. import util 29 30 31URL = { 32 "images": { 33 "train": "http://resource.deepwise.com/ChestX-Det/train_data.zip", 34 "test": "http://resource.deepwise.com/ChestX-Det/test_data.zip", 35 }, 36 "annotations": { 37 "train": "https://raw.githubusercontent.com/Deepwise-AILab/ChestX-Det-Dataset/main/ChestX_Det_train.json", 38 "test": "https://raw.githubusercontent.com/Deepwise-AILab/ChestX-Det-Dataset/main/ChestX_Det_test.json", 39 }, 40} 41 42CHECKSUM = { 43 "images": { 44 "train": "413f74a03383280e2d63f6215c8eb581aa386cfec323fb4148f0916d2f5f2900", 45 "test": "c52677d1e4043bf425bc997d62db2a780f6260ea1d67ad37c65fe3b2ffdef14f", 46 }, 47 "annotations": { 48 "train": None, 49 "test": None, 50 }, 51} 52 53CHESTX_DET_LABELS = { 54 0: "background", 55 1: "Atelectasis", 56 2: "Calcification", 57 3: "Cardiomegaly", 58 4: "Consolidation", 59 5: "Diffuse Nodule", 60 6: "Effusion", 61 7: "Emphysema", 62 8: "Fibrosis", 63 9: "Fracture", 64 10: "Mass", 65 11: "Nodule", 66 12: "Pleural Thickening", 67 13: "Pneumothorax", 68} 69"""The label ids of the 13 categories of thoracic abnormalities or diseases annotated in ChestX-Det.""" 70 71LABEL_IDS = {name: label_id for label_id, name in CHESTX_DET_LABELS.items() if label_id != 0} 72 73 74def get_chestx_det_data(path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False) -> str: 75 """Download the ChestX-Det data. 76 77 Args: 78 path: Filepath to a folder where the data is downloaded for further processing. 79 split: The choice of data split. 80 download: Whether to download the data if it is not present. 81 82 Returns: 83 Filepath where the image data is downloaded. 84 """ 85 if split not in ("train", "test"): 86 raise ValueError(f"'{split}' is not a valid split.") 87 88 image_dir = os.path.join(path, split) 89 if os.path.exists(image_dir): 90 return image_dir 91 92 os.makedirs(path, exist_ok=True) 93 94 zip_path = os.path.join(path, f"{split}_data.zip") 95 util.download_source(path=zip_path, url=URL["images"][split], download=download, checksum=CHECKSUM["images"][split]) 96 util.unzip(zip_path=zip_path, dst=path, remove=False) 97 98 annotation_path = os.path.join(path, f"ChestX_Det_{split}.json") 99 util.download_source( 100 path=annotation_path, url=URL["annotations"][split], download=download, checksum=CHECKSUM["annotations"][split] 101 ) 102 103 return image_dir 104 105 106def _rasterize_annotations(shape, syms, polygons): 107 labels = np.zeros(shape, dtype="uint8") 108 for sym, poly in zip(syms, polygons): 109 poly = np.asarray(poly) 110 rr, cc = polygon(poly[:, 1], poly[:, 0], shape=shape) 111 labels[rr, cc] = LABEL_IDS[sym] 112 return labels 113 114 115def _preprocess_split(image_dir, annotation_path, preprocessed_dir): 116 os.makedirs(preprocessed_dir, exist_ok=True) 117 118 with open(annotation_path) as f: 119 annotations = json.load(f) 120 121 image_paths, gt_paths = [], [] 122 for ann in tqdm(annotations, desc=f"Preprocessing labels for {image_dir}"): 123 image_path = os.path.join(image_dir, ann["file_name"]) 124 if not os.path.exists(image_path): 125 continue 126 127 gt_path = os.path.join(preprocessed_dir, ann["file_name"]) 128 if not os.path.exists(gt_path): 129 shape = imageio.imread(image_path).shape[:2] 130 labels = _rasterize_annotations(shape, ann["syms"], ann["polygons"]) 131 imageio.imwrite(gt_path, labels) 132 133 image_paths.append(image_path) 134 gt_paths.append(gt_path) 135 136 return image_paths, gt_paths 137 138 139def get_chestx_det_paths( 140 path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False 141) -> Tuple[List[str], List[str]]: 142 """Get paths to the ChestX-Det data. 143 144 Args: 145 path: Filepath to a folder where the data is downloaded for further processing. 146 split: The choice of data split. 147 download: Whether to download the data if it is not present. 148 149 Returns: 150 List of filepaths for the image data. 151 List of filepaths for the label data. 152 """ 153 image_dir = get_chestx_det_data(path=path, split=split, download=download) 154 155 annotation_path = os.path.join(path, f"ChestX_Det_{split}.json") 156 preprocessed_dir = os.path.join(path, "preprocessed", split) 157 158 if os.path.exists(preprocessed_dir) and len(glob(os.path.join(preprocessed_dir, "*.png"))) > 0: 159 image_paths = natsorted(glob(os.path.join(image_dir, "*.png"))) 160 gt_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.png"))) 161 return image_paths, gt_paths 162 163 image_paths, gt_paths = _preprocess_split(image_dir, annotation_path, preprocessed_dir) 164 return natsorted(image_paths), natsorted(gt_paths) 165 166 167def get_chestx_det_dataset( 168 path: Union[os.PathLike, str], 169 patch_shape: Tuple[int, int], 170 split: Literal["train", "test"], 171 resize_inputs: bool = False, 172 download: bool = False, 173 **kwargs 174) -> Dataset: 175 """Get the ChestX-Det dataset for segmentation of thoracic abnormalities in chest x-rays. 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 split: The choice of data split. 181 resize_inputs: Whether to resize the inputs to the expected 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, gt_paths = get_chestx_det_paths(path, split, 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=None, 199 label_paths=gt_paths, 200 label_key=None, 201 patch_shape=patch_shape, 202 is_seg_dataset=False, 203 **kwargs 204 ) 205 206 207def get_chestx_det_loader( 208 path: Union[os.PathLike, str], 209 batch_size: int, 210 patch_shape: Tuple[int, int], 211 split: Literal["train", "test"], 212 resize_inputs: bool = False, 213 download: bool = False, 214 **kwargs 215) -> DataLoader: 216 """Get the ChestX-Det dataloader for segmentation of thoracic abnormalities in chest x-rays. 217 218 Args: 219 path: Filepath to a folder where the data is downloaded for further processing. 220 batch_size: The batch size for training. 221 patch_shape: The patch shape to use for training. 222 split: The choice of data split. 223 resize_inputs: Whether to resize the inputs to the expected patch shape. 224 download: Whether to download the data if it is not present. 225 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 226 227 Returns: 228 The DataLoader. 229 """ 230 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 231 dataset = get_chestx_det_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 232 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
The label ids of the 13 categories of thoracic abnormalities or diseases annotated in ChestX-Det.
75def get_chestx_det_data(path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False) -> str: 76 """Download the ChestX-Det data. 77 78 Args: 79 path: Filepath to a folder where the data is downloaded for further processing. 80 split: The choice of data split. 81 download: Whether to download the data if it is not present. 82 83 Returns: 84 Filepath where the image data is downloaded. 85 """ 86 if split not in ("train", "test"): 87 raise ValueError(f"'{split}' is not a valid split.") 88 89 image_dir = os.path.join(path, split) 90 if os.path.exists(image_dir): 91 return image_dir 92 93 os.makedirs(path, exist_ok=True) 94 95 zip_path = os.path.join(path, f"{split}_data.zip") 96 util.download_source(path=zip_path, url=URL["images"][split], download=download, checksum=CHECKSUM["images"][split]) 97 util.unzip(zip_path=zip_path, dst=path, remove=False) 98 99 annotation_path = os.path.join(path, f"ChestX_Det_{split}.json") 100 util.download_source( 101 path=annotation_path, url=URL["annotations"][split], download=download, checksum=CHECKSUM["annotations"][split] 102 ) 103 104 return image_dir
Download the ChestX-Det data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The choice of data split.
- download: Whether to download the data if it is not present.
Returns:
Filepath where the image data is downloaded.
140def get_chestx_det_paths( 141 path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False 142) -> Tuple[List[str], List[str]]: 143 """Get paths to the ChestX-Det data. 144 145 Args: 146 path: Filepath to a folder where the data is downloaded for further processing. 147 split: The choice of data split. 148 download: Whether to download the data if it is not present. 149 150 Returns: 151 List of filepaths for the image data. 152 List of filepaths for the label data. 153 """ 154 image_dir = get_chestx_det_data(path=path, split=split, download=download) 155 156 annotation_path = os.path.join(path, f"ChestX_Det_{split}.json") 157 preprocessed_dir = os.path.join(path, "preprocessed", split) 158 159 if os.path.exists(preprocessed_dir) and len(glob(os.path.join(preprocessed_dir, "*.png"))) > 0: 160 image_paths = natsorted(glob(os.path.join(image_dir, "*.png"))) 161 gt_paths = natsorted(glob(os.path.join(preprocessed_dir, "*.png"))) 162 return image_paths, gt_paths 163 164 image_paths, gt_paths = _preprocess_split(image_dir, annotation_path, preprocessed_dir) 165 return natsorted(image_paths), natsorted(gt_paths)
Get paths to the ChestX-Det data.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- split: The choice of data split.
- 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.
168def get_chestx_det_dataset( 169 path: Union[os.PathLike, str], 170 patch_shape: Tuple[int, int], 171 split: Literal["train", "test"], 172 resize_inputs: bool = False, 173 download: bool = False, 174 **kwargs 175) -> Dataset: 176 """Get the ChestX-Det dataset for segmentation of thoracic abnormalities in chest x-rays. 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 split: The choice of data split. 182 resize_inputs: Whether to resize the inputs to the expected patch shape. 183 download: Whether to download the data if it is not present. 184 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`. 185 186 Returns: 187 The segmentation dataset. 188 """ 189 image_paths, gt_paths = get_chestx_det_paths(path, split, download) 190 191 if resize_inputs: 192 resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False} 193 kwargs, patch_shape = util.update_kwargs_for_resize_trafo( 194 kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs 195 ) 196 197 return torch_em.default_segmentation_dataset( 198 raw_paths=image_paths, 199 raw_key=None, 200 label_paths=gt_paths, 201 label_key=None, 202 patch_shape=patch_shape, 203 is_seg_dataset=False, 204 **kwargs 205 )
Get the ChestX-Det dataset for segmentation of thoracic abnormalities in chest x-rays.
Arguments:
- path: Filepath to a folder where the data is downloaded for further processing.
- patch_shape: The patch shape to use for training.
- split: The choice of data split.
- resize_inputs: Whether to resize the inputs to the expected 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.
208def get_chestx_det_loader( 209 path: Union[os.PathLike, str], 210 batch_size: int, 211 patch_shape: Tuple[int, int], 212 split: Literal["train", "test"], 213 resize_inputs: bool = False, 214 download: bool = False, 215 **kwargs 216) -> DataLoader: 217 """Get the ChestX-Det dataloader for segmentation of thoracic abnormalities in chest x-rays. 218 219 Args: 220 path: Filepath to a folder where the data is downloaded for further processing. 221 batch_size: The batch size for training. 222 patch_shape: The patch shape to use for training. 223 split: The choice of data split. 224 resize_inputs: Whether to resize the inputs to the expected patch shape. 225 download: Whether to download the data if it is not present. 226 kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader. 227 228 Returns: 229 The DataLoader. 230 """ 231 ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs) 232 dataset = get_chestx_det_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs) 233 return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
Get the ChestX-Det dataloader for segmentation of thoracic abnormalities in chest x-rays.
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.
- split: The choice of data split.
- resize_inputs: Whether to resize the inputs to the expected 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.