torch_em.data.datasets.medical.fracatlas

The FracAtlas dataset contains annotations for fracture segmentation in musculoskeletal radiographs.

The dataset consists of 4,083 radiographs of the hand, leg, hip and shoulder, of which 719 fractured images come with fracture segmentation polygons (COCO annotations with the single category 'fractured', up to several polygons per image). The remaining 3,366 images are non-fractured and have no annotations, so this loader only exposes the 719 annotated images. The polygons are rasterized into binary masks (1 = fracture) during preprocessing, and the images are stored as single-channel tif files, since the JPEGs are grayscale but a part of them is saved with three identical channels. The image sizes vary from 454x373 to 2880x2304 pixels, use resize_inputs=True to train with batches.

The official splits ('train', 'val' and 'test', from 'Utilities/Fracture Split') are available via the split argument. They cover exactly the annotated images.

The dataset is located at https://doi.org/10.6084/m9.figshare.22363012, released under a CC-BY-4.0 license.

This dataset is from the publication https://doi.org/10.1038/s41597-023-02432-4. Please cite it if you use this dataset for your research.

  1"""The FracAtlas dataset contains annotations for fracture segmentation in musculoskeletal radiographs.
  2
  3The dataset consists of 4,083 radiographs of the hand, leg, hip and shoulder, of which 719 fractured
  4images come with fracture segmentation polygons (COCO annotations with the single category 'fractured', up to
  5several polygons per image). The remaining 3,366 images are non-fractured and have no annotations, so this
  6loader only exposes the 719 annotated images. The polygons are rasterized into binary masks
  7(1 = fracture) during preprocessing, and the images are stored as single-channel tif files, since the JPEGs
  8are grayscale but a part of them is saved with three identical channels. The image sizes vary from
  9454x373 to 2880x2304 pixels, use `resize_inputs=True` to train with batches.
 10
 11The official splits ('train', 'val' and 'test', from 'Utilities/Fracture Split') are available via the
 12`split` argument. They cover exactly the annotated images.
 13
 14The dataset is located at https://doi.org/10.6084/m9.figshare.22363012, released under a CC-BY-4.0 license.
 15
 16This dataset is from the publication https://doi.org/10.1038/s41597-023-02432-4.
 17Please cite it if you use this dataset for your research.
 18"""
 19
 20import os
 21import json
 22import uuid
 23from tqdm import tqdm
 24from typing import Union, Tuple, Literal, List
 25
 26import numpy as np
 27import imageio.v3 as imageio
 28
 29from torch.utils.data import Dataset, DataLoader
 30
 31import torch_em
 32
 33from .. import util
 34
 35
 36URL = "https://ndownloader.figshare.com/files/65518038"
 37CHECKSUM = "b67ec2d290a022b3dcf47f78e9a37f7edcc80592c0571f439355bf00bd9f0e23"
 38
 39SPLITS = ["train", "val", "test"]
 40SPLIT_FILES = {"train": "train.csv", "val": "valid.csv", "test": "test.csv"}
 41
 42
 43def _write_atomic(path, array):
 44    tmp_path = f"{path}.{uuid.uuid4().hex}.incomplete.tif"
 45    imageio.imwrite(tmp_path, array, extension=".tif")
 46    os.replace(tmp_path, path)
 47
 48
 49def _preprocess_data(data_dir, preprocessed_dir):
 50    from skimage.draw import polygon
 51
 52    with open(os.path.join(data_dir, "Annotations", "COCO JSON", "COCO_fracture_masks.json")) as f:
 53        coco = json.load(f)
 54
 55    image_dir = os.path.join(preprocessed_dir, "images")
 56    label_dir = os.path.join(preprocessed_dir, "labels")
 57    os.makedirs(image_dir, exist_ok=True)
 58    os.makedirs(label_dir, exist_ok=True)
 59
 60    polygons = {}
 61    for annotation in coco["annotations"]:
 62        polygons.setdefault(annotation["image_id"], []).extend(annotation["segmentation"])
 63
 64    for image_info in tqdm(coco["images"], desc="Preprocess FracAtlas"):
 65        if image_info["id"] not in polygons:
 66            continue
 67
 68        stem = os.path.splitext(image_info["file_name"])[0]
 69        image_path = os.path.join(image_dir, f"{stem}.tif")
 70        label_path = os.path.join(label_dir, f"{stem}.tif")
 71        if os.path.exists(image_path) and os.path.exists(label_path):
 72            continue
 73
 74        image = imageio.imread(os.path.join(data_dir, "images", "Fractured", image_info["file_name"]))
 75        if image.ndim == 3:
 76            image = image[..., 0]
 77
 78        label = np.zeros(image.shape, dtype="uint8")
 79        for coordinates in polygons[image_info["id"]]:
 80            coordinates = np.asarray(coordinates, dtype="float64").reshape(-1, 2)
 81            rr, cc = polygon(coordinates[:, 1], coordinates[:, 0], shape=label.shape)
 82            label[rr, cc] = 1
 83
 84        _write_atomic(label_path, label)
 85        _write_atomic(image_path, image)
 86
 87
 88def get_fracatlas_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 89    """Download the FracAtlas dataset and rasterize the fracture polygons into masks.
 90
 91    Args:
 92        path: Filepath to a folder where the data is downloaded for further processing.
 93        download: Whether to download the data if it is not present.
 94
 95    Returns:
 96        Filepath to the extracted dataset.
 97    """
 98    data_dir = os.path.join(path, "FracAtlas")
 99    if not os.path.exists(data_dir):
100        os.makedirs(path, exist_ok=True)
101        zip_path = os.path.join(path, "FracAtlas.zip")
102        util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
103        util.unzip(zip_path=zip_path, dst=path, remove=False)
104
105    _preprocess_data(data_dir, os.path.join(path, "preprocessed"))
106    return data_dir
107
108
109def get_fracatlas_paths(
110    path: Union[os.PathLike, str], split: Literal["train", "val", "test"], download: bool = False,
111) -> Tuple[List[str], List[str]]:
112    """Get paths to the FracAtlas data.
113
114    Args:
115        path: Filepath to a folder where the data is downloaded for further processing.
116        split: The choice of data split. Either 'train', 'val' or 'test'.
117        download: Whether to download the data if it is not present.
118
119    Returns:
120        List of filepaths for the image data.
121        List of filepaths for the label data.
122    """
123    if split not in SPLITS:
124        raise ValueError(f"'{split}' is not a valid split. Choose one of {SPLITS}.")
125
126    data_dir = get_fracatlas_data(path, download)
127
128    with open(os.path.join(data_dir, "Utilities", "Fracture Split", SPLIT_FILES[split])) as f:
129        names = [os.path.splitext(line.strip())[0] for line in f.read().splitlines()[1:] if line.strip()]
130
131    preprocessed_dir = os.path.join(path, "preprocessed")
132    raw_paths = [os.path.join(preprocessed_dir, "images", f"{name}.tif") for name in sorted(names)]
133    label_paths = [os.path.join(preprocessed_dir, "labels", f"{name}.tif") for name in sorted(names)]
134
135    assert len(raw_paths) > 0
136    assert all(os.path.exists(p) for p in raw_paths + label_paths)
137
138    return raw_paths, label_paths
139
140
141def get_fracatlas_dataset(
142    path: Union[os.PathLike, str],
143    patch_shape: Tuple[int, int],
144    split: Literal["train", "val", "test"],
145    resize_inputs: bool = False,
146    download: bool = False,
147    **kwargs
148) -> Dataset:
149    """Get the FracAtlas dataset for fracture segmentation in musculoskeletal radiographs.
150
151    Args:
152        path: Filepath to a folder where the data is downloaded for further processing.
153        patch_shape: The patch shape to use for training.
154        split: The choice of data split. Either 'train', 'val' or 'test'.
155        resize_inputs: Whether to resize the inputs to the patch shape.
156        download: Whether to download the data if it is not present.
157        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
158
159    Returns:
160        The segmentation dataset.
161    """
162    raw_paths, label_paths = get_fracatlas_paths(path, split, download)
163
164    if resize_inputs:
165        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
166        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
167            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
168        )
169
170    return torch_em.default_segmentation_dataset(
171        raw_paths=raw_paths,
172        raw_key=None,
173        label_paths=label_paths,
174        label_key=None,
175        is_seg_dataset=False,
176        patch_shape=patch_shape,
177        **kwargs
178    )
179
180
181def get_fracatlas_loader(
182    path: Union[os.PathLike, str],
183    batch_size: int,
184    patch_shape: Tuple[int, int],
185    split: Literal["train", "val", "test"],
186    resize_inputs: bool = False,
187    download: bool = False,
188    **kwargs
189) -> DataLoader:
190    """Get the FracAtlas dataloader for fracture segmentation in musculoskeletal radiographs.
191
192    Args:
193        path: Filepath to a folder where the data is downloaded for further processing.
194        batch_size: The batch size for training.
195        patch_shape: The patch shape to use for training.
196        split: The choice of data split. Either 'train', 'val' or 'test'.
197        resize_inputs: Whether to resize the inputs to the patch shape.
198        download: Whether to download the data if it is not present.
199        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
200
201    Returns:
202        The DataLoader.
203    """
204    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
205    dataset = get_fracatlas_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
206    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
URL = 'https://ndownloader.figshare.com/files/65518038'
CHECKSUM = 'b67ec2d290a022b3dcf47f78e9a37f7edcc80592c0571f439355bf00bd9f0e23'
SPLITS = ['train', 'val', 'test']
SPLIT_FILES = {'train': 'train.csv', 'val': 'valid.csv', 'test': 'test.csv'}
def get_fracatlas_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 89def get_fracatlas_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 90    """Download the FracAtlas dataset and rasterize the fracture polygons into masks.
 91
 92    Args:
 93        path: Filepath to a folder where the data is downloaded for further processing.
 94        download: Whether to download the data if it is not present.
 95
 96    Returns:
 97        Filepath to the extracted dataset.
 98    """
 99    data_dir = os.path.join(path, "FracAtlas")
100    if not os.path.exists(data_dir):
101        os.makedirs(path, exist_ok=True)
102        zip_path = os.path.join(path, "FracAtlas.zip")
103        util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
104        util.unzip(zip_path=zip_path, dst=path, remove=False)
105
106    _preprocess_data(data_dir, os.path.join(path, "preprocessed"))
107    return data_dir

Download the FracAtlas dataset and rasterize the fracture polygons into masks.

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 to the extracted dataset.

def get_fracatlas_paths( path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False) -> Tuple[List[str], List[str]]:
110def get_fracatlas_paths(
111    path: Union[os.PathLike, str], split: Literal["train", "val", "test"], download: bool = False,
112) -> Tuple[List[str], List[str]]:
113    """Get paths to the FracAtlas data.
114
115    Args:
116        path: Filepath to a folder where the data is downloaded for further processing.
117        split: The choice of data split. Either 'train', 'val' or 'test'.
118        download: Whether to download the data if it is not present.
119
120    Returns:
121        List of filepaths for the image data.
122        List of filepaths for the label data.
123    """
124    if split not in SPLITS:
125        raise ValueError(f"'{split}' is not a valid split. Choose one of {SPLITS}.")
126
127    data_dir = get_fracatlas_data(path, download)
128
129    with open(os.path.join(data_dir, "Utilities", "Fracture Split", SPLIT_FILES[split])) as f:
130        names = [os.path.splitext(line.strip())[0] for line in f.read().splitlines()[1:] if line.strip()]
131
132    preprocessed_dir = os.path.join(path, "preprocessed")
133    raw_paths = [os.path.join(preprocessed_dir, "images", f"{name}.tif") for name in sorted(names)]
134    label_paths = [os.path.join(preprocessed_dir, "labels", f"{name}.tif") for name in sorted(names)]
135
136    assert len(raw_paths) > 0
137    assert all(os.path.exists(p) for p in raw_paths + label_paths)
138
139    return raw_paths, label_paths

Get paths to the FracAtlas data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split. Either 'train', 'val' or 'test'.
  • 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_fracatlas_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'val', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
142def get_fracatlas_dataset(
143    path: Union[os.PathLike, str],
144    patch_shape: Tuple[int, int],
145    split: Literal["train", "val", "test"],
146    resize_inputs: bool = False,
147    download: bool = False,
148    **kwargs
149) -> Dataset:
150    """Get the FracAtlas dataset for fracture segmentation in musculoskeletal radiographs.
151
152    Args:
153        path: Filepath to a folder where the data is downloaded for further processing.
154        patch_shape: The patch shape to use for training.
155        split: The choice of data split. Either 'train', 'val' or 'test'.
156        resize_inputs: Whether to resize the inputs to the patch shape.
157        download: Whether to download the data if it is not present.
158        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
159
160    Returns:
161        The segmentation dataset.
162    """
163    raw_paths, label_paths = get_fracatlas_paths(path, split, download)
164
165    if resize_inputs:
166        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
167        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
168            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
169        )
170
171    return torch_em.default_segmentation_dataset(
172        raw_paths=raw_paths,
173        raw_key=None,
174        label_paths=label_paths,
175        label_key=None,
176        is_seg_dataset=False,
177        patch_shape=patch_shape,
178        **kwargs
179    )

Get the FracAtlas dataset for fracture segmentation in musculoskeletal radiographs.

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. Either 'train', 'val' or 'test'.
  • 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_fracatlas_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, int], split: Literal['train', 'val', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
182def get_fracatlas_loader(
183    path: Union[os.PathLike, str],
184    batch_size: int,
185    patch_shape: Tuple[int, int],
186    split: Literal["train", "val", "test"],
187    resize_inputs: bool = False,
188    download: bool = False,
189    **kwargs
190) -> DataLoader:
191    """Get the FracAtlas dataloader for fracture segmentation in musculoskeletal radiographs.
192
193    Args:
194        path: Filepath to a folder where the data is downloaded for further processing.
195        batch_size: The batch size for training.
196        patch_shape: The patch shape to use for training.
197        split: The choice of data split. Either 'train', 'val' or 'test'.
198        resize_inputs: Whether to resize the inputs to the patch shape.
199        download: Whether to download the data if it is not present.
200        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
201
202    Returns:
203        The DataLoader.
204    """
205    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
206    dataset = get_fracatlas_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
207    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the FracAtlas dataloader for fracture segmentation in musculoskeletal 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.
  • split: The choice of data split. Either 'train', 'val' or 'test'.
  • 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.