torch_em.data.datasets.medical.focus

FOCUS is the Four-chamber Ultrasound Image Dataset for Fetal Cardiac Biometric Measurement, with annotations for fetal heart and thorax segmentation in four-chamber view ultrasound images, used e.g. to estimate the cardiothoracic diameter ratio.

The dataset is located at https://zenodo.org/records/14597550 (CC BY 4.0). This dataset is from Zenodo, with DOI https://doi.org/10.5281/zenodo.14597550. Please cite it if you use this dataset for your research.

  1"""FOCUS is the Four-chamber Ultrasound Image Dataset for Fetal Cardiac Biometric
  2Measurement, with annotations for fetal heart and thorax segmentation in four-chamber
  3view ultrasound images, used e.g. to estimate the cardiothoracic diameter ratio.
  4
  5The dataset is located at https://zenodo.org/records/14597550 (CC BY 4.0).
  6This dataset is from Zenodo, with DOI https://doi.org/10.5281/zenodo.14597550.
  7Please cite it if you use this dataset for your research.
  8"""
  9
 10import os
 11from glob import glob
 12from tqdm import tqdm
 13from typing import Union, Tuple, List, Literal
 14
 15import numpy as np
 16import imageio.v3 as imageio
 17
 18from torch.utils.data import Dataset, DataLoader
 19
 20import torch_em
 21
 22from .. import util
 23
 24
 25URL = "https://zenodo.org/records/14597550/files/FOCUS-dataset.zip"
 26CHECKSUM = "625ee59d9d8adfb946790f03bd04e6342ad2a12499a0bed3ec02b65ec35369b8"
 27
 28SPLIT_FOLDERS = {"train": "training", "val": "validation", "test": "testing"}
 29
 30
 31def get_focus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 32    """Download the FOCUS dataset.
 33
 34    Args:
 35        path: Filepath to a folder where the data is downloaded for further processing.
 36        download: Whether to download the data if it is not present.
 37
 38    Returns:
 39        Filepath to the folder with the downloaded images and segmentation masks.
 40    """
 41    if os.path.exists(os.path.join(path, "training")):
 42        return path
 43
 44    os.makedirs(path, exist_ok=True)
 45
 46    zip_path = os.path.join(path, "focus.zip")
 47    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
 48    util.unzip(zip_path=zip_path, dst=path)
 49
 50    return path
 51
 52
 53def get_focus_paths(
 54    path: Union[os.PathLike, str], split: Literal["train", "val", "test"], download: bool = False
 55) -> Tuple[List[str], List[str]]:
 56    """Get paths to the FOCUS data.
 57
 58    Args:
 59        path: Filepath to a folder where the data is downloaded for further processing.
 60        split: The choice of data split.
 61        download: Whether to download the data if it is not present.
 62
 63    Returns:
 64        List of filepaths for the image data.
 65        List of filepaths for the label data.
 66    """
 67    if split not in SPLIT_FOLDERS:
 68        raise ValueError(f"'{split}' is not a supported split. Choose one of {list(SPLIT_FOLDERS.keys())}.")
 69
 70    data_dir = get_focus_data(path=path, download=download)
 71    split_dir = os.path.join(data_dir, SPLIT_FOLDERS[split])
 72
 73    image_paths = sorted(glob(os.path.join(split_dir, "images", "*.png")))
 74
 75    label_dir = os.path.join(split_dir, "annfiles_semantic")
 76    os.makedirs(label_dir, exist_ok=True)
 77
 78    gt_paths = []
 79    for image_path in tqdm(image_paths, desc=f"Preprocessing FOCUS '{split}' labels"):
 80        fname = os.path.splitext(os.path.basename(image_path))[0]
 81        gt_path = os.path.join(label_dir, f"{fname}.tif")
 82        gt_paths.append(gt_path)
 83        if os.path.exists(gt_path):
 84            continue
 85
 86        thorax = imageio.imread(os.path.join(split_dir, "annfiles_mask", f"{fname}-thorax.png"))
 87        cardiac = imageio.imread(os.path.join(split_dir, "annfiles_mask", f"{fname}-cardiac.png"))
 88        if thorax.ndim == 3:
 89            thorax = thorax[..., 0]
 90        if cardiac.ndim == 3:
 91            cardiac = cardiac[..., 0]
 92
 93        label = np.zeros(thorax.shape, dtype="uint8")
 94        label[thorax > 127] = 1
 95        label[cardiac > 127] = 2
 96
 97        imageio.imwrite(gt_path, label, compression="zlib")
 98
 99    return image_paths, gt_paths
100
101
102def get_focus_dataset(
103    path: Union[os.PathLike, str],
104    patch_shape: Tuple[int, int],
105    split: Literal["train", "val", "test"],
106    resize_inputs: bool = False,
107    download: bool = False,
108    **kwargs
109) -> Dataset:
110    """Get the FOCUS dataset for fetal cardiac and thorax segmentation.
111
112    Args:
113        path: Filepath to a folder where the data is downloaded for further processing.
114        patch_shape: The patch shape to use for training.
115        split: The choice of data split.
116        resize_inputs: Whether to resize the inputs to the patch shape.
117        download: Whether to download the data if it is not present.
118        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
119
120    Returns:
121        The segmentation dataset.
122    """
123    image_paths, gt_paths = get_focus_paths(path, split, download)
124
125    if resize_inputs:
126        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
127        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
128            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
129        )
130
131    return torch_em.default_segmentation_dataset(
132        raw_paths=image_paths,
133        raw_key=None,
134        label_paths=gt_paths,
135        label_key=None,
136        patch_shape=patch_shape,
137        is_seg_dataset=False,
138        **kwargs
139    )
140
141
142def get_focus_loader(
143    path: Union[os.PathLike, str],
144    patch_shape: Tuple[int, int],
145    batch_size: int,
146    split: Literal["train", "val", "test"],
147    resize_inputs: bool = False,
148    download: bool = False,
149    **kwargs
150) -> DataLoader:
151    """Get the FOCUS dataloader for fetal cardiac and thorax segmentation.
152
153    Args:
154        path: Filepath to a folder where the data is downloaded for further processing.
155        patch_shape: The patch shape to use for training.
156        batch_size: The batch size for training.
157        split: The choice of data split.
158        resize_inputs: Whether to resize the inputs to the patch shape.
159        download: Whether to download the data if it is not present.
160        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
161
162    Returns:
163        The DataLoader.
164    """
165    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
166    dataset = get_focus_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
167    return torch_em.get_data_loader(dataset=dataset, batch_size=batch_size, **loader_kwargs)
URL = 'https://zenodo.org/records/14597550/files/FOCUS-dataset.zip'
CHECKSUM = '625ee59d9d8adfb946790f03bd04e6342ad2a12499a0bed3ec02b65ec35369b8'
SPLIT_FOLDERS = {'train': 'training', 'val': 'validation', 'test': 'testing'}
def get_focus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
32def get_focus_data(path: Union[os.PathLike, str], download: bool = False) -> str:
33    """Download the FOCUS dataset.
34
35    Args:
36        path: Filepath to a folder where the data is downloaded for further processing.
37        download: Whether to download the data if it is not present.
38
39    Returns:
40        Filepath to the folder with the downloaded images and segmentation masks.
41    """
42    if os.path.exists(os.path.join(path, "training")):
43        return path
44
45    os.makedirs(path, exist_ok=True)
46
47    zip_path = os.path.join(path, "focus.zip")
48    util.download_source(path=zip_path, url=URL, download=download, checksum=CHECKSUM)
49    util.unzip(zip_path=zip_path, dst=path)
50
51    return path

Download the FOCUS 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 to the folder with the downloaded images and segmentation masks.

def get_focus_paths( path: Union[os.PathLike, str], split: Literal['train', 'val', 'test'], download: bool = False) -> Tuple[List[str], List[str]]:
 54def get_focus_paths(
 55    path: Union[os.PathLike, str], split: Literal["train", "val", "test"], download: bool = False
 56) -> Tuple[List[str], List[str]]:
 57    """Get paths to the FOCUS data.
 58
 59    Args:
 60        path: Filepath to a folder where the data is downloaded for further processing.
 61        split: The choice of data split.
 62        download: Whether to download the data if it is not present.
 63
 64    Returns:
 65        List of filepaths for the image data.
 66        List of filepaths for the label data.
 67    """
 68    if split not in SPLIT_FOLDERS:
 69        raise ValueError(f"'{split}' is not a supported split. Choose one of {list(SPLIT_FOLDERS.keys())}.")
 70
 71    data_dir = get_focus_data(path=path, download=download)
 72    split_dir = os.path.join(data_dir, SPLIT_FOLDERS[split])
 73
 74    image_paths = sorted(glob(os.path.join(split_dir, "images", "*.png")))
 75
 76    label_dir = os.path.join(split_dir, "annfiles_semantic")
 77    os.makedirs(label_dir, exist_ok=True)
 78
 79    gt_paths = []
 80    for image_path in tqdm(image_paths, desc=f"Preprocessing FOCUS '{split}' labels"):
 81        fname = os.path.splitext(os.path.basename(image_path))[0]
 82        gt_path = os.path.join(label_dir, f"{fname}.tif")
 83        gt_paths.append(gt_path)
 84        if os.path.exists(gt_path):
 85            continue
 86
 87        thorax = imageio.imread(os.path.join(split_dir, "annfiles_mask", f"{fname}-thorax.png"))
 88        cardiac = imageio.imread(os.path.join(split_dir, "annfiles_mask", f"{fname}-cardiac.png"))
 89        if thorax.ndim == 3:
 90            thorax = thorax[..., 0]
 91        if cardiac.ndim == 3:
 92            cardiac = cardiac[..., 0]
 93
 94        label = np.zeros(thorax.shape, dtype="uint8")
 95        label[thorax > 127] = 1
 96        label[cardiac > 127] = 2
 97
 98        imageio.imwrite(gt_path, label, compression="zlib")
 99
100    return image_paths, gt_paths

Get paths to the FOCUS 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_focus_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:
103def get_focus_dataset(
104    path: Union[os.PathLike, str],
105    patch_shape: Tuple[int, int],
106    split: Literal["train", "val", "test"],
107    resize_inputs: bool = False,
108    download: bool = False,
109    **kwargs
110) -> Dataset:
111    """Get the FOCUS dataset for fetal cardiac and thorax segmentation.
112
113    Args:
114        path: Filepath to a folder where the data is downloaded for further processing.
115        patch_shape: The patch shape to use for training.
116        split: The choice of data split.
117        resize_inputs: Whether to resize the inputs to the patch shape.
118        download: Whether to download the data if it is not present.
119        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
120
121    Returns:
122        The segmentation dataset.
123    """
124    image_paths, gt_paths = get_focus_paths(path, split, download)
125
126    if resize_inputs:
127        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
128        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
129            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
130        )
131
132    return torch_em.default_segmentation_dataset(
133        raw_paths=image_paths,
134        raw_key=None,
135        label_paths=gt_paths,
136        label_key=None,
137        patch_shape=patch_shape,
138        is_seg_dataset=False,
139        **kwargs
140    )

Get the FOCUS dataset for fetal cardiac and thorax segmentation.

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 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_focus_loader( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], batch_size: int, split: Literal['train', 'val', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
143def get_focus_loader(
144    path: Union[os.PathLike, str],
145    patch_shape: Tuple[int, int],
146    batch_size: int,
147    split: Literal["train", "val", "test"],
148    resize_inputs: bool = False,
149    download: bool = False,
150    **kwargs
151) -> DataLoader:
152    """Get the FOCUS dataloader for fetal cardiac and thorax segmentation.
153
154    Args:
155        path: Filepath to a folder where the data is downloaded for further processing.
156        patch_shape: The patch shape to use for training.
157        batch_size: The batch size for training.
158        split: The choice of data split.
159        resize_inputs: Whether to resize the inputs to the patch shape.
160        download: Whether to download the data if it is not present.
161        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
162
163    Returns:
164        The DataLoader.
165    """
166    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
167    dataset = get_focus_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
168    return torch_em.get_data_loader(dataset=dataset, batch_size=batch_size, **loader_kwargs)

Get the FOCUS dataloader for fetal cardiac and thorax 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.
  • split: The choice of data split.
  • 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.