torch_em.data.datasets.medical.bagls

BAGLS is a benchmark dataset for glottis segmentation in high-speed videolaryngoscopy recordings, collected across 7 institutions. Each frame comes with a pixel-wise binary segmentation of the glottal area (not every frame has an annotated glottis).

The dataset is located at https://doi.org/10.5281/zenodo.3762320. This dataset is from the publication https://doi.org/10.1038/s41597-020-0526-3. Please cite it if you use this dataset for your research.

  1"""BAGLS is a benchmark dataset for glottis segmentation in high-speed videolaryngoscopy
  2recordings, collected across 7 institutions. Each frame comes with a pixel-wise binary
  3segmentation of the glottal area (not every frame has an annotated glottis).
  4
  5The dataset is located at https://doi.org/10.5281/zenodo.3762320.
  6This dataset is from the publication https://doi.org/10.1038/s41597-020-0526-3.
  7Please cite it if you use this dataset for your research.
  8"""
  9
 10import os
 11from glob import glob
 12from typing import Union, Tuple, List, Literal
 13
 14from torch.utils.data import Dataset, DataLoader
 15
 16import torch_em
 17
 18from .. import util
 19
 20
 21URLS = {
 22    "train": "https://zenodo.org/records/3762320/files/training.zip",
 23    "test": "https://zenodo.org/records/3762320/files/test.zip",
 24}
 25
 26CHECKSUMS = {
 27    "train": "7850fd4666131b2d6d5e6bbc544a2609955195b936d9a3d9f34c3e6f67c24b1c",
 28    "test": "f41ea31bacb46d2a2924f149249a5a9419e7533720dba9b5b86eb5b1d1f984b9",
 29}
 30
 31
 32def get_bagls_data(path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False) -> str:
 33    """Download the BAGLS dataset.
 34
 35    Args:
 36        path: Filepath to a folder where the data is downloaded for further processing.
 37        split: The choice of data split.
 38        download: Whether to download the data if it is not present.
 39
 40    Returns:
 41        Filepath where the data is downloaded.
 42    """
 43    data_dir = os.path.join(path, split)
 44    if os.path.exists(data_dir):
 45        return data_dir
 46
 47    os.makedirs(path, exist_ok=True)
 48
 49    zip_name = "training.zip" if split == "train" else "test.zip"
 50    zip_path = os.path.join(path, zip_name)
 51    util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split])
 52    util.unzip(zip_path=zip_path, dst=path)
 53
 54    return data_dir
 55
 56
 57def get_bagls_paths(
 58    path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False
 59) -> Tuple[List[str], List[str]]:
 60    """Get paths to the BAGLS data.
 61
 62    Args:
 63        path: Filepath to a folder where the data is downloaded for further processing.
 64        split: The choice of data split.
 65        download: Whether to download the data if it is not present.
 66
 67    Returns:
 68        List of filepaths for the image data.
 69        List of filepaths for the label data.
 70    """
 71    data_dir = get_bagls_data(path=path, split=split, download=download)
 72
 73    gt_paths = sorted(glob(os.path.join(data_dir, "*_seg.png")))
 74    image_paths = [gt_path.replace("_seg.png", ".png") for gt_path in gt_paths]
 75
 76    return image_paths, gt_paths
 77
 78
 79def get_bagls_dataset(
 80    path: Union[os.PathLike, str],
 81    patch_shape: Tuple[int, int],
 82    split: Literal["train", "test"],
 83    resize_inputs: bool = False,
 84    download: bool = False,
 85    **kwargs
 86) -> Dataset:
 87    """Get the BAGLS dataset for glottis segmentation.
 88
 89    Args:
 90        path: Filepath to a folder where the data is downloaded for further processing.
 91        patch_shape: The patch shape to use for training.
 92        split: The choice of data split.
 93        resize_inputs: Whether to resize the inputs to the patch shape.
 94        download: Whether to download the data if it is not present.
 95        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
 96
 97    Returns:
 98        The segmentation dataset.
 99    """
100    image_paths, gt_paths = get_bagls_paths(path, split, download)
101
102    if resize_inputs:
103        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
104        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
105            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
106        )
107
108    return torch_em.default_segmentation_dataset(
109        raw_paths=image_paths,
110        raw_key=None,
111        label_paths=gt_paths,
112        label_key=None,
113        patch_shape=patch_shape,
114        is_seg_dataset=False,
115        **kwargs
116    )
117
118
119def get_bagls_loader(
120    path: Union[os.PathLike, str],
121    patch_shape: Tuple[int, int],
122    batch_size: int,
123    split: Literal["train", "test"],
124    resize_inputs: bool = False,
125    download: bool = False,
126    **kwargs
127) -> DataLoader:
128    """Get the BAGLS dataloader for glottis 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        split: The choice of data split.
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_bagls_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
144    return torch_em.get_data_loader(dataset=dataset, batch_size=batch_size, **loader_kwargs)
URLS = {'train': 'https://zenodo.org/records/3762320/files/training.zip', 'test': 'https://zenodo.org/records/3762320/files/test.zip'}
CHECKSUMS = {'train': '7850fd4666131b2d6d5e6bbc544a2609955195b936d9a3d9f34c3e6f67c24b1c', 'test': 'f41ea31bacb46d2a2924f149249a5a9419e7533720dba9b5b86eb5b1d1f984b9'}
def get_bagls_data( path: Union[os.PathLike, str], split: Literal['train', 'test'], download: bool = False) -> str:
33def get_bagls_data(path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False) -> str:
34    """Download the BAGLS dataset.
35
36    Args:
37        path: Filepath to a folder where the data is downloaded for further processing.
38        split: The choice of data split.
39        download: Whether to download the data if it is not present.
40
41    Returns:
42        Filepath where the data is downloaded.
43    """
44    data_dir = os.path.join(path, split)
45    if os.path.exists(data_dir):
46        return data_dir
47
48    os.makedirs(path, exist_ok=True)
49
50    zip_name = "training.zip" if split == "train" else "test.zip"
51    zip_path = os.path.join(path, zip_name)
52    util.download_source(path=zip_path, url=URLS[split], download=download, checksum=CHECKSUMS[split])
53    util.unzip(zip_path=zip_path, dst=path)
54
55    return data_dir

Download the BAGLS dataset.

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:

Filepath where the data is downloaded.

def get_bagls_paths( path: Union[os.PathLike, str], split: Literal['train', 'test'], download: bool = False) -> Tuple[List[str], List[str]]:
58def get_bagls_paths(
59    path: Union[os.PathLike, str], split: Literal["train", "test"], download: bool = False
60) -> Tuple[List[str], List[str]]:
61    """Get paths to the BAGLS data.
62
63    Args:
64        path: Filepath to a folder where the data is downloaded for further processing.
65        split: The choice of data split.
66        download: Whether to download the data if it is not present.
67
68    Returns:
69        List of filepaths for the image data.
70        List of filepaths for the label data.
71    """
72    data_dir = get_bagls_data(path=path, split=split, download=download)
73
74    gt_paths = sorted(glob(os.path.join(data_dir, "*_seg.png")))
75    image_paths = [gt_path.replace("_seg.png", ".png") for gt_path in gt_paths]
76
77    return image_paths, gt_paths

Get paths to the BAGLS 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_bagls_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
 80def get_bagls_dataset(
 81    path: Union[os.PathLike, str],
 82    patch_shape: Tuple[int, int],
 83    split: Literal["train", "test"],
 84    resize_inputs: bool = False,
 85    download: bool = False,
 86    **kwargs
 87) -> Dataset:
 88    """Get the BAGLS dataset for glottis segmentation.
 89
 90    Args:
 91        path: Filepath to a folder where the data is downloaded for further processing.
 92        patch_shape: The patch shape to use for training.
 93        split: The choice of data split.
 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_bagls_paths(path, split, download)
102
103    if resize_inputs:
104        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": True}
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    )

Get the BAGLS dataset for glottis 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_bagls_loader( path: Union[os.PathLike, str], patch_shape: Tuple[int, int], batch_size: int, split: Literal['train', 'test'], resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
120def get_bagls_loader(
121    path: Union[os.PathLike, str],
122    patch_shape: Tuple[int, int],
123    batch_size: int,
124    split: Literal["train", "test"],
125    resize_inputs: bool = False,
126    download: bool = False,
127    **kwargs
128) -> DataLoader:
129    """Get the BAGLS dataloader for glottis 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        split: The choice of data split.
136        resize_inputs: Whether to resize the inputs to the patch shape.
137        download: Whether to download the data if it is not present.
138        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
139
140    Returns:
141        The DataLoader.
142    """
143    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
144    dataset = get_bagls_dataset(path, patch_shape, split, resize_inputs, download, **ds_kwargs)
145    return torch_em.get_data_loader(dataset=dataset, batch_size=batch_size, **loader_kwargs)

Get the BAGLS dataloader for glottis 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.