torch_em.data.datasets.medical.chest_xray_masks

The Chest Xray Masks and Labels dataset contains annotations for lung segmentation in chest x-ray images.

The dataset combines the Montgomery County and Shenzhen chest x-ray collections with manually drawn lung field masks and is located at https://www.kaggle.com/datasets/nikhilpandey360/chest-xray-masks-and-labels. The underlying data is released by the National Library of Medicine, National Institutes of Health, Bethesda, MD, USA and Shenzhen No.3 People's Hospital, Guangdong Medical College, Shenzhen, China. This dataset is from the publications https://doi.org/10.1109/TMI.2013.2284099 and https://doi.org/10.1109/TMI.2013.2290491. Please cite them if you use this dataset for your research.

  1"""The Chest Xray Masks and Labels dataset contains annotations for lung
  2segmentation in chest x-ray images.
  3
  4The dataset combines the Montgomery County and Shenzhen chest x-ray collections
  5with manually drawn lung field masks and is located at
  6https://www.kaggle.com/datasets/nikhilpandey360/chest-xray-masks-and-labels.
  7The underlying data is released by the National Library of Medicine, National
  8Institutes of Health, Bethesda, MD, USA and Shenzhen No.3 People's Hospital,
  9Guangdong Medical College, Shenzhen, China. This dataset is from the publications
 10https://doi.org/10.1109/TMI.2013.2284099 and https://doi.org/10.1109/TMI.2013.2290491.
 11Please cite them if you use this dataset for your research.
 12"""
 13
 14import os
 15from glob import glob
 16from natsort import natsorted
 17from typing import Union, Tuple, List
 18
 19from torch.utils.data import Dataset, DataLoader
 20
 21import torch_em
 22
 23from .. import util
 24
 25
 26KAGGLE_DATASET_NAME = "nikhilpandey360/chest-xray-masks-and-labels"
 27
 28
 29def get_chest_xray_masks_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 30    """Download the Chest Xray Masks and Labels dataset.
 31
 32    Args:
 33        path: Filepath to a folder where the data is downloaded for further processing.
 34        download: Whether to download the data if it is not present.
 35
 36    Returns:
 37        Filepath where the data is downloaded.
 38    """
 39    data_dir = os.path.join(path, "Lung Segmentation")
 40    if os.path.exists(data_dir):
 41        return data_dir
 42
 43    os.makedirs(path, exist_ok=True)
 44
 45    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET_NAME, download=download)
 46    zip_path = os.path.join(path, "chest-xray-masks-and-labels.zip")
 47    util.unzip(zip_path=zip_path, dst=path)
 48
 49    return data_dir
 50
 51
 52def get_chest_xray_masks_paths(
 53    path: Union[os.PathLike, str], download: bool = False
 54) -> Tuple[List[str], List[str]]:
 55    """Get paths to the Chest Xray Masks and Labels data.
 56
 57    Args:
 58        path: Filepath to a folder where the data is downloaded for further processing.
 59        download: Whether to download the data if it is not present.
 60
 61    Returns:
 62        List of filepaths for the image data.
 63        List of filepaths for the label data.
 64    """
 65    data_dir = get_chest_xray_masks_data(path=path, download=download)
 66
 67    mask_paths = natsorted(glob(os.path.join(data_dir, "masks", "*.png")))
 68
 69    # NOTE: Masks for the Shenzhen images have a '_mask' suffix (eg. 'CHNCXR_0001_0_mask.png'), while masks
 70    # for the Montgomery images share the same filename as the corresponding image (eg. 'MCUCXR_0001_0.png').
 71    image_paths = []
 72    for mask_path in mask_paths:
 73        image_id = os.path.basename(mask_path).replace("_mask.png", ".png")
 74        image_paths.append(os.path.join(data_dir, "CXR_png", image_id))
 75
 76    assert len(image_paths) == len(mask_paths) and len(image_paths) > 0
 77    assert all(os.path.exists(p) for p in image_paths)
 78
 79    return image_paths, mask_paths
 80
 81
 82def get_chest_xray_masks_dataset(
 83    path: Union[os.PathLike, str],
 84    patch_shape: Tuple[int, int],
 85    resize_inputs: bool = False,
 86    download: bool = False,
 87    **kwargs
 88) -> Dataset:
 89    """Get the Chest Xray Masks and Labels dataset for lung segmentation.
 90
 91    Args:
 92        path: Filepath to a folder where the data is downloaded for further processing.
 93        patch_shape: The patch shape to use for training.
 94        resize_inputs: Whether to resize the inputs to the patch shape.
 95        download: Whether to download the data if it is not present.
 96        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
 97
 98    Returns:
 99        The segmentation dataset.
100    """
101    image_paths, gt_paths = get_chest_xray_masks_paths(path, download)
102
103    if resize_inputs:
104        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
105        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
106            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
107        )
108
109    return torch_em.default_segmentation_dataset(
110        raw_paths=image_paths,
111        raw_key=None,
112        label_paths=gt_paths,
113        label_key=None,
114        patch_shape=patch_shape,
115        is_seg_dataset=False,
116        **kwargs
117    )
118
119
120def get_chest_xray_masks_loader(
121    path: Union[os.PathLike, str],
122    patch_shape: Tuple[int, int],
123    batch_size: int,
124    resize_inputs: bool = False,
125    download: bool = False,
126    **kwargs
127) -> DataLoader:
128    """Get the Chest Xray Masks and Labels dataloader for lung segmentation.
129
130    Args:
131        path: Filepath to a folder where the data is downloaded for further processing.
132        patch_shape: The patch shape to use for training.
133        batch_size: The batch size for training.
134        resize_inputs: Whether to resize the inputs to the patch shape.
135        download: Whether to download the data if it is not present.
136        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
137
138    Returns:
139        The DataLoader.
140    """
141    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
142    dataset = get_chest_xray_masks_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
143    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
KAGGLE_DATASET_NAME = 'nikhilpandey360/chest-xray-masks-and-labels'
def get_chest_xray_masks_data(path: Union[os.PathLike, str], download: bool = False) -> str:
30def get_chest_xray_masks_data(path: Union[os.PathLike, str], download: bool = False) -> str:
31    """Download the Chest Xray Masks and Labels dataset.
32
33    Args:
34        path: Filepath to a folder where the data is downloaded for further processing.
35        download: Whether to download the data if it is not present.
36
37    Returns:
38        Filepath where the data is downloaded.
39    """
40    data_dir = os.path.join(path, "Lung Segmentation")
41    if os.path.exists(data_dir):
42        return data_dir
43
44    os.makedirs(path, exist_ok=True)
45
46    util.download_source_kaggle(path=path, dataset_name=KAGGLE_DATASET_NAME, download=download)
47    zip_path = os.path.join(path, "chest-xray-masks-and-labels.zip")
48    util.unzip(zip_path=zip_path, dst=path)
49
50    return data_dir

Download the Chest Xray Masks and Labels 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 data is downloaded.

def get_chest_xray_masks_paths( path: Union[os.PathLike, str], download: bool = False) -> Tuple[List[str], List[str]]:
53def get_chest_xray_masks_paths(
54    path: Union[os.PathLike, str], download: bool = False
55) -> Tuple[List[str], List[str]]:
56    """Get paths to the Chest Xray Masks and Labels data.
57
58    Args:
59        path: Filepath to a folder where the data is downloaded for further processing.
60        download: Whether to download the data if it is not present.
61
62    Returns:
63        List of filepaths for the image data.
64        List of filepaths for the label data.
65    """
66    data_dir = get_chest_xray_masks_data(path=path, download=download)
67
68    mask_paths = natsorted(glob(os.path.join(data_dir, "masks", "*.png")))
69
70    # NOTE: Masks for the Shenzhen images have a '_mask' suffix (eg. 'CHNCXR_0001_0_mask.png'), while masks
71    # for the Montgomery images share the same filename as the corresponding image (eg. 'MCUCXR_0001_0.png').
72    image_paths = []
73    for mask_path in mask_paths:
74        image_id = os.path.basename(mask_path).replace("_mask.png", ".png")
75        image_paths.append(os.path.join(data_dir, "CXR_png", image_id))
76
77    assert len(image_paths) == len(mask_paths) and len(image_paths) > 0
78    assert all(os.path.exists(p) for p in image_paths)
79
80    return image_paths, mask_paths

Get paths to the Chest Xray Masks and Labels 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_chest_xray_masks_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
 83def get_chest_xray_masks_dataset(
 84    path: Union[os.PathLike, str],
 85    patch_shape: Tuple[int, int],
 86    resize_inputs: bool = False,
 87    download: bool = False,
 88    **kwargs
 89) -> Dataset:
 90    """Get the Chest Xray Masks and Labels dataset for lung segmentation.
 91
 92    Args:
 93        path: Filepath to a folder where the data is downloaded for further processing.
 94        patch_shape: The patch shape to use for training.
 95        resize_inputs: Whether to resize the inputs to the patch shape.
 96        download: Whether to download the data if it is not present.
 97        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
 98
 99    Returns:
100        The segmentation dataset.
101    """
102    image_paths, gt_paths = get_chest_xray_masks_paths(path, download)
103
104    if resize_inputs:
105        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
106        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
107            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
108        )
109
110    return torch_em.default_segmentation_dataset(
111        raw_paths=image_paths,
112        raw_key=None,
113        label_paths=gt_paths,
114        label_key=None,
115        patch_shape=patch_shape,
116        is_seg_dataset=False,
117        **kwargs
118    )

Get the Chest Xray Masks and Labels dataset for lung 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 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_chest_xray_masks_loader( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], batch_size: int, resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
121def get_chest_xray_masks_loader(
122    path: Union[os.PathLike, str],
123    patch_shape: Tuple[int, int],
124    batch_size: int,
125    resize_inputs: bool = False,
126    download: bool = False,
127    **kwargs
128) -> DataLoader:
129    """Get the Chest Xray Masks and Labels dataloader for lung segmentation.
130
131    Args:
132        path: Filepath to a folder where the data is downloaded for further processing.
133        patch_shape: The patch shape to use for training.
134        batch_size: The batch size for training.
135        resize_inputs: Whether to resize the inputs to the patch shape.
136        download: Whether to download the data if it is not present.
137        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
138
139    Returns:
140        The DataLoader.
141    """
142    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
143    dataset = get_chest_xray_masks_dataset(path, patch_shape, resize_inputs, download, **ds_kwargs)
144    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the Chest Xray Masks and Labels dataloader for lung segmentation.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • batch_size: The batch size 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.