torch_em.data.datasets.medical.crosspan

CrossPan is a benchmark for pancreas segmentation in MRI, across MRI sequences and institutions.

It consists of 1336 3D MRI volumes from 8 institutions across three sequences: T1-weighted (T1), T2-weighted (T2) and Out-of-Phase (OOP), each pre-split into 'train', 'val' and 'test' subsets, with binary pancreas segmentation masks.

NOTE: The label legend is as follows: background: 0, pancreas: 1. Verified on the data: the label volumes only contain the ids 0 and 1.

The dataset is located at https://huggingface.co/datasets/linkai-peng/CrossPan (CC BY-NC 4.0). This dataset is from the publication https://doi.org/10.48550/arXiv.2604.18797. Please cite it if you use this dataset in your research.

  1"""CrossPan is a benchmark for pancreas segmentation in MRI, across MRI sequences and institutions.
  2
  3It consists of 1336 3D MRI volumes from 8 institutions across three sequences: T1-weighted (T1), T2-weighted
  4(T2) and Out-of-Phase (OOP), each pre-split into 'train', 'val' and 'test' subsets, with binary pancreas
  5segmentation masks.
  6
  7NOTE: The label legend is as follows: background: 0, pancreas: 1. Verified on the data: the label
  8volumes only contain the ids 0 and 1.
  9
 10The dataset is located at https://huggingface.co/datasets/linkai-peng/CrossPan (CC BY-NC 4.0).
 11This dataset is from the publication https://doi.org/10.48550/arXiv.2604.18797.
 12Please cite it if you use this dataset in your research.
 13"""
 14
 15import os
 16from glob import glob
 17from natsort import natsorted
 18from typing import Union, Tuple, List, Literal, Optional
 19
 20from torch.utils.data import Dataset, DataLoader
 21
 22import torch_em
 23
 24from .. import util
 25
 26
 27HF_REPO = "linkai-peng/CrossPan"
 28
 29SEQUENCES = ["T1", "T2", "OOP"]
 30SPLITS = ["train", "val", "test"]
 31
 32LABEL_IDS = {"background": 0, "pancreas": 1}
 33
 34
 35def get_crosspan_data(path: Union[os.PathLike, str], download: bool = False) -> str:
 36    """Download the CrossPan dataset.
 37
 38    Args:
 39        path: Filepath to a folder where the data is downloaded for further processing.
 40        download: Whether to download the data if it is not present.
 41
 42    Returns:
 43        Filepath where the data is stored.
 44    """
 45    if os.path.exists(path):
 46        return path
 47
 48    if not download:
 49        raise RuntimeError(f"Cannot find the data at '{path}', but download was set to False.")
 50
 51    try:
 52        from huggingface_hub import snapshot_download
 53    except ImportError:
 54        raise ImportError("'huggingface_hub' is required to download CrossPan. Install it via conda/pip.")
 55
 56    os.makedirs(path, exist_ok=True)
 57    snapshot_download(repo_id=HF_REPO, repo_type="dataset", local_dir=path)
 58
 59    return path
 60
 61
 62def get_crosspan_paths(
 63    path: Union[os.PathLike, str],
 64    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
 65    split: Optional[Literal["train", "val", "test"]] = None,
 66    download: bool = False,
 67) -> Tuple[List[str], List[str]]:
 68    """Get paths to the CrossPan data.
 69
 70    Args:
 71        path: Filepath to a folder where the data is downloaded for further processing.
 72        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
 73        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
 74        download: Whether to download the data if it is not present.
 75
 76    Returns:
 77        List of filepaths for the image data.
 78        List of filepaths for the label data.
 79    """
 80    data_dir = get_crosspan_data(path, download)
 81
 82    if sequence is None:
 83        sequences = SEQUENCES
 84    elif sequence in SEQUENCES:
 85        sequences = [sequence]
 86    else:
 87        raise ValueError(f"'{sequence}' is not a valid sequence.")
 88
 89    if split is None:
 90        splits = SPLITS
 91    elif split in SPLITS:
 92        splits = [split]
 93    else:
 94        raise ValueError(f"'{split}' is not a valid split.")
 95
 96    raw_paths, label_paths = [], []
 97    for seq in sequences:
 98        for spl in splits:
 99            cur_raw_paths = natsorted(glob(os.path.join(data_dir, seq, spl, "images", "*_0000.nii.gz")))
100            cur_label_paths = [
101                p.replace(f"{os.sep}images{os.sep}", f"{os.sep}labels{os.sep}").replace("_0000.nii.gz", ".nii.gz")
102                for p in cur_raw_paths
103            ]
104            assert len(cur_raw_paths) > 0 and all(os.path.exists(p) for p in cur_label_paths)
105            raw_paths.extend(cur_raw_paths)
106            label_paths.extend(cur_label_paths)
107
108    return raw_paths, label_paths
109
110
111def get_crosspan_dataset(
112    path: Union[os.PathLike, str],
113    patch_shape: Tuple[int, ...],
114    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
115    split: Optional[Literal["train", "val", "test"]] = None,
116    resize_inputs: bool = False,
117    download: bool = False,
118    **kwargs
119) -> Dataset:
120    """Get the CrossPan dataset for pancreas segmentation in MRI.
121
122    Args:
123        path: Filepath to a folder where the data is downloaded for further processing.
124        patch_shape: The patch shape to use for training.
125        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
126        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
127        resize_inputs: Whether to resize inputs to the desired patch shape.
128        download: Whether to download the data if it is not present.
129        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
130
131    Returns:
132        The segmentation dataset.
133    """
134    raw_paths, label_paths = get_crosspan_paths(path, sequence, split, download)
135
136    if resize_inputs:
137        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
138        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
139            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
140        )
141
142    return torch_em.default_segmentation_dataset(
143        raw_paths=raw_paths,
144        raw_key="data",
145        label_paths=label_paths,
146        label_key="data",
147        patch_shape=patch_shape,
148        is_seg_dataset=True,
149        **kwargs
150    )
151
152
153def get_crosspan_loader(
154    path: Union[os.PathLike, str],
155    batch_size: int,
156    patch_shape: Tuple[int, ...],
157    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
158    split: Optional[Literal["train", "val", "test"]] = None,
159    resize_inputs: bool = False,
160    download: bool = False,
161    **kwargs
162) -> DataLoader:
163    """Get the CrossPan dataloader for pancreas segmentation in MRI.
164
165    Args:
166        path: Filepath to a folder where the data is downloaded for further processing.
167        batch_size: The batch size for training.
168        patch_shape: The patch shape to use for training.
169        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
170        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
171        resize_inputs: Whether to resize inputs to the desired patch shape.
172        download: Whether to download the data if it is not present.
173        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
174
175    Returns:
176        The DataLoader.
177    """
178    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
179    dataset = get_crosspan_dataset(path, patch_shape, sequence, split, resize_inputs, download, **ds_kwargs)
180    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
HF_REPO = 'linkai-peng/CrossPan'
SEQUENCES = ['T1', 'T2', 'OOP']
SPLITS = ['train', 'val', 'test']
LABEL_IDS = {'background': 0, 'pancreas': 1}
def get_crosspan_data(path: Union[os.PathLike, str], download: bool = False) -> str:
36def get_crosspan_data(path: Union[os.PathLike, str], download: bool = False) -> str:
37    """Download the CrossPan dataset.
38
39    Args:
40        path: Filepath to a folder where the data is downloaded for further processing.
41        download: Whether to download the data if it is not present.
42
43    Returns:
44        Filepath where the data is stored.
45    """
46    if os.path.exists(path):
47        return path
48
49    if not download:
50        raise RuntimeError(f"Cannot find the data at '{path}', but download was set to False.")
51
52    try:
53        from huggingface_hub import snapshot_download
54    except ImportError:
55        raise ImportError("'huggingface_hub' is required to download CrossPan. Install it via conda/pip.")
56
57    os.makedirs(path, exist_ok=True)
58    snapshot_download(repo_id=HF_REPO, repo_type="dataset", local_dir=path)
59
60    return path

Download the CrossPan 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 stored.

def get_crosspan_paths( path: Union[os.PathLike, str], sequence: Optional[Literal['T1', 'T2', 'OOP']] = None, split: Optional[Literal['train', 'val', 'test']] = None, download: bool = False) -> Tuple[List[str], List[str]]:
 63def get_crosspan_paths(
 64    path: Union[os.PathLike, str],
 65    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
 66    split: Optional[Literal["train", "val", "test"]] = None,
 67    download: bool = False,
 68) -> Tuple[List[str], List[str]]:
 69    """Get paths to the CrossPan data.
 70
 71    Args:
 72        path: Filepath to a folder where the data is downloaded for further processing.
 73        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
 74        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
 75        download: Whether to download the data if it is not present.
 76
 77    Returns:
 78        List of filepaths for the image data.
 79        List of filepaths for the label data.
 80    """
 81    data_dir = get_crosspan_data(path, download)
 82
 83    if sequence is None:
 84        sequences = SEQUENCES
 85    elif sequence in SEQUENCES:
 86        sequences = [sequence]
 87    else:
 88        raise ValueError(f"'{sequence}' is not a valid sequence.")
 89
 90    if split is None:
 91        splits = SPLITS
 92    elif split in SPLITS:
 93        splits = [split]
 94    else:
 95        raise ValueError(f"'{split}' is not a valid split.")
 96
 97    raw_paths, label_paths = [], []
 98    for seq in sequences:
 99        for spl in splits:
100            cur_raw_paths = natsorted(glob(os.path.join(data_dir, seq, spl, "images", "*_0000.nii.gz")))
101            cur_label_paths = [
102                p.replace(f"{os.sep}images{os.sep}", f"{os.sep}labels{os.sep}").replace("_0000.nii.gz", ".nii.gz")
103                for p in cur_raw_paths
104            ]
105            assert len(cur_raw_paths) > 0 and all(os.path.exists(p) for p in cur_label_paths)
106            raw_paths.extend(cur_raw_paths)
107            label_paths.extend(cur_label_paths)
108
109    return raw_paths, label_paths

Get paths to the CrossPan data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
  • split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
  • 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_crosspan_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], sequence: Optional[Literal['T1', 'T2', 'OOP']] = None, split: Optional[Literal['train', 'val', 'test']] = None, resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
112def get_crosspan_dataset(
113    path: Union[os.PathLike, str],
114    patch_shape: Tuple[int, ...],
115    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
116    split: Optional[Literal["train", "val", "test"]] = None,
117    resize_inputs: bool = False,
118    download: bool = False,
119    **kwargs
120) -> Dataset:
121    """Get the CrossPan dataset for pancreas segmentation in MRI.
122
123    Args:
124        path: Filepath to a folder where the data is downloaded for further processing.
125        patch_shape: The patch shape to use for training.
126        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
127        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
128        resize_inputs: Whether to resize inputs to the desired patch shape.
129        download: Whether to download the data if it is not present.
130        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
131
132    Returns:
133        The segmentation dataset.
134    """
135    raw_paths, label_paths = get_crosspan_paths(path, sequence, split, download)
136
137    if resize_inputs:
138        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
139        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
140            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
141        )
142
143    return torch_em.default_segmentation_dataset(
144        raw_paths=raw_paths,
145        raw_key="data",
146        label_paths=label_paths,
147        label_key="data",
148        patch_shape=patch_shape,
149        is_seg_dataset=True,
150        **kwargs
151    )

Get the CrossPan dataset for pancreas segmentation in MRI.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • patch_shape: The patch shape to use for training.
  • sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
  • split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
  • resize_inputs: Whether to resize inputs to the desired 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_crosspan_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], sequence: Optional[Literal['T1', 'T2', 'OOP']] = None, split: Optional[Literal['train', 'val', 'test']] = None, resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
154def get_crosspan_loader(
155    path: Union[os.PathLike, str],
156    batch_size: int,
157    patch_shape: Tuple[int, ...],
158    sequence: Optional[Literal["T1", "T2", "OOP"]] = None,
159    split: Optional[Literal["train", "val", "test"]] = None,
160    resize_inputs: bool = False,
161    download: bool = False,
162    **kwargs
163) -> DataLoader:
164    """Get the CrossPan dataloader for pancreas segmentation in MRI.
165
166    Args:
167        path: Filepath to a folder where the data is downloaded for further processing.
168        batch_size: The batch size for training.
169        patch_shape: The patch shape to use for training.
170        sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
171        split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
172        resize_inputs: Whether to resize inputs to the desired patch shape.
173        download: Whether to download the data if it is not present.
174        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or for the PyTorch DataLoader.
175
176    Returns:
177        The DataLoader.
178    """
179    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
180    dataset = get_crosspan_dataset(path, patch_shape, sequence, split, resize_inputs, download, **ds_kwargs)
181    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the CrossPan dataloader for pancreas segmentation in MRI.

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.
  • sequence: The choice of MRI sequence. Either 'T1', 'T2' or 'OOP'. If None, all sequences are used.
  • split: The choice of data split. Either 'train', 'val' or 'test'. If None, all splits are used.
  • resize_inputs: Whether to resize inputs to the desired 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.