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)
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.
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.
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.
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_datasetor for the PyTorch DataLoader.
Returns:
The DataLoader.