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 ('_frontal.png' or '_lateral.png') paired with a boolean multi-channel label volume of shape (159, 512, 512) ('labels/_.npy'), one channel per anatomical class. Of the 159 classes, 51 are 'direct' (directly annotated structures) and 108 are 'derived' (hierarchical unions of the direct structures, e.g. 'spine' is the union of 'cervical spine', 'thoracic spine' and 'lumbar spine', which are themselves unions of individual vertebrae). 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)
URLS = {'images': 'https://zenodo.org/records/20688438/files/paxray_images.tar.gz', 'labels': 'https://zenodo.org/records/20688438/files/paxray_labels.zip', 'label_dictionary': 'https://zenodo.org/records/20688438/files/paxray_label_dictionary.json', 'label_types': 'https://zenodo.org/records/20688438/files/paxray_label_types.json'}
CHECKSUMS = {'images': 'e91a9d4d27b5bbfbd24046c897b7facf5391ab934936559fd75b1c92e988b99f', 'labels': '9dc4623db0be9793be781699af0605f26a764ab5088efffbb9395adc0e2082c1', 'label_dictionary': 'f5bf6f48117f18e9cfbc2f844ad7369d043926e32fc4eeae1638b36b05da39a1', 'label_types': '465282f18b246a9b8631a145cf7a8d4ccdc4841acea2cc7eb3eedaa4ce633768'}
def get_paxray_pp_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 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.

def get_paxray_pp_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
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.

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

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

The DataLoader.