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