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)
URL = {'images': {'train': 'http://resource.deepwise.com/ChestX-Det/train_data.zip', 'test': 'http://resource.deepwise.com/ChestX-Det/test_data.zip'}, 'annotations': {'train': 'https://raw.githubusercontent.com/Deepwise-AILab/ChestX-Det-Dataset/main/ChestX_Det_train.json', 'test': 'https://raw.githubusercontent.com/Deepwise-AILab/ChestX-Det-Dataset/main/ChestX_Det_test.json'}}
CHECKSUM = {'images': {'train': '413f74a03383280e2d63f6215c8eb581aa386cfec323fb4148f0916d2f5f2900', 'test': 'c52677d1e4043bf425bc997d62db2a780f6260ea1d67ad37c65fe3b2ffdef14f'}, 'annotations': {'train': None, 'test': None}}
CHESTX_DET_LABELS = {0: 'background', 1: 'Atelectasis', 2: 'Calcification', 3: 'Cardiomegaly', 4: 'Consolidation', 5: 'Diffuse Nodule', 6: 'Effusion', 7: 'Emphysema', 8: 'Fibrosis', 9: 'Fracture', 10: 'Mass', 11: 'Nodule', 12: 'Pleural Thickening', 13: 'Pneumothorax'}

The label ids of the 13 categories of thoracic abnormalities or diseases annotated in ChestX-Det.

LABEL_IDS = {'Atelectasis': 1, 'Calcification': 2, 'Cardiomegaly': 3, 'Consolidation': 4, 'Diffuse Nodule': 5, 'Effusion': 6, 'Emphysema': 7, 'Fibrosis': 8, 'Fracture': 9, 'Mass': 10, 'Nodule': 11, 'Pleural Thickening': 12, 'Pneumothorax': 13}
def get_chestx_det_data( path: Union[os.PathLike, str], split: Literal['train', 'test'], download: bool = False) -> str:
 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.

def get_chestx_det_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'], download: bool = False) -> Tuple[List[str], List[str]]:
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.

def get_chestx_det_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
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.

def get_chestx_det_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
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_dataset or for the PyTorch DataLoader.
Returns:

The DataLoader.