torch_em.data.datasets.medical.rexgroundingct

ReXGroundingCT links free-text chest radiology findings to pixel-level 3D lesion / finding segmentations in non-contrast chest CT scans.

The dataset re-uses the CT volumes from CT-RATE (https://doi.org/10.48550/arXiv.2403.17834) and adds segmentation masks for 8,028 findings across 14 abnormality categories in 3,142 CT scans: 2,992 scans for training, 50 for public validation, and 100 held out privately for the MICCAI 2026 challenge leaderboard. Each raw mask file is a 4D volume of shape (finding_category, X, Y, Z), one channel per abnormality category present in that scan; within a channel, distinct positive values label distinct entities of that finding (see anatomical_cot.json / dataset.json on the dataset repository for the full per-finding metadata). This loader merges all categories/entities of a scan into a single 3D instance segmentation volume with a globally unique instance id per (category, entity) pair.

NOTE: The masks are hosted at https://huggingface.co/datasets/rajpurkarlab/ReXGroundingCT, but the CT volumes themselves are NOT included in that repository. They have to be fetched separately from CT-RATE at https://huggingface.co/datasets/ibrahimhamamci/CT-RATE, using the matching case names (e.g. the mask 'segmentations/train_10000_a_1.nii.gz' corresponds to the CT-RATE volume at 'dataset/train_fixed/train_10000/train_10000_a/train_10000_a_1.nii.gz'). Both repositories are gated on HuggingFace: visiting the dataset pages and accepting the license terms while logged in (self-service, no manual review) is required before downloading with a HuggingFace access token.

NOTE: This loader only exposes the paired CT volume and finding-instance mask as a plain segmentation dataset. It does not parse or expose the free-text finding descriptions / categories from 'dataset.json' or 'reports_dataset.json'; use get_rexgroundingct_metadata to load the raw per-case metadata dictionary from 'dataset.json' if the associated text is needed.

The masks are licensed under CC BY-NC-SA 4.0. This dataset is from the publication https://doi.org/10.48550/arXiv.2507.22030. Please cite it if you use this dataset in your research.

  1"""ReXGroundingCT links free-text chest radiology findings to pixel-level 3D lesion / finding
  2segmentations in non-contrast chest CT scans.
  3
  4The dataset re-uses the CT volumes from CT-RATE (https://doi.org/10.48550/arXiv.2403.17834) and adds
  5segmentation masks for 8,028 findings across 14 abnormality categories in 3,142 CT scans: 2,992 scans
  6for training, 50 for public validation, and 100 held out privately for the MICCAI 2026 challenge
  7leaderboard. Each raw mask file is a 4D volume of shape (finding_category, X, Y, Z), one channel per
  8abnormality category present in that scan; within a channel, distinct positive values label distinct
  9entities of that finding (see `anatomical_cot.json` / `dataset.json` on the dataset repository for the
 10full per-finding metadata). This loader merges all categories/entities of a scan into a single 3D
 11instance segmentation volume with a globally unique instance id per (category, entity) pair.
 12
 13NOTE: The masks are hosted at https://huggingface.co/datasets/rajpurkarlab/ReXGroundingCT, but the CT
 14volumes themselves are NOT included in that repository. They have to be fetched separately from CT-RATE
 15at https://huggingface.co/datasets/ibrahimhamamci/CT-RATE, using the matching case names (e.g. the mask
 16'segmentations/train_10000_a_1.nii.gz' corresponds to the CT-RATE volume at
 17'dataset/train_fixed/train_10000/train_10000_a/train_10000_a_1.nii.gz'). Both repositories are gated on
 18HuggingFace: visiting the dataset pages and accepting the license terms while logged in (self-service,
 19no manual review) is required before downloading with a HuggingFace access token.
 20
 21NOTE: This loader only exposes the paired CT volume and finding-instance mask as a plain segmentation
 22dataset. It does not parse or expose the free-text finding descriptions / categories from
 23'dataset.json' or 'reports_dataset.json'; use `get_rexgroundingct_metadata` to load the raw per-case
 24metadata dictionary from 'dataset.json' if the associated text is needed.
 25
 26The masks are licensed under CC BY-NC-SA 4.0. This dataset is from the publication
 27https://doi.org/10.48550/arXiv.2507.22030. Please cite it if you use this dataset in your research.
 28"""
 29
 30import os
 31import json
 32from glob import glob
 33from natsort import natsorted
 34from typing import Union, Tuple, List, Literal, Optional, Dict, Any
 35
 36import numpy as np
 37
 38from torch.utils.data import Dataset, DataLoader
 39
 40import torch_em
 41
 42from .. import util
 43
 44
 45REXGROUNDINGCT_REPO = "rajpurkarlab/ReXGroundingCT"
 46CT_RATE_REPO = "ibrahimhamamci/CT-RATE"
 47
 48SPLITS = ["train", "val"]
 49
 50SPLIT_TAGS = {"train": "train", "val": "valid"}
 51"""Mapping from the split name used by this module to the case name prefix used on the dataset repos."""
 52
 53
 54def _ct_rate_path(name: str) -> str:
 55    # e.g. 'train_10000_a_1' -> 'dataset/train_fixed/train_10000/train_10000_a/train_10000_a_1.nii.gz'
 56    parts = name.split("_")
 57    if len(parts) != 4:
 58        raise ValueError(f"Unexpected ReXGroundingCT case name format: '{name}'")
 59    tag, case_id, letter, _idx = parts
 60    case_dir = f"{tag}_{case_id}"
 61    recon_dir = f"{tag}_{case_id}_{letter}"
 62    return f"dataset/{tag}_fixed/{case_dir}/{recon_dir}/{name}.nii.gz"
 63
 64
 65def _case_names_for_split(dataset_json: Dict[str, Any], split: str) -> List[str]:
 66    tag = SPLIT_TAGS[split]
 67
 68    if split in dataset_json and isinstance(dataset_json[split], (list, dict)):
 69        entries = dataset_json[split]
 70        names = list(entries.keys()) if isinstance(entries, dict) else [e["name"] for e in entries]
 71    elif all(isinstance(v, dict) and "name" in v for v in dataset_json.values()):
 72        names = [v["name"] for v in dataset_json.values() if v["name"].startswith(f"{tag}_")]
 73    elif isinstance(dataset_json, dict) and all(
 74        k.startswith((SPLIT_TAGS["train"], SPLIT_TAGS["val"])) for k in dataset_json
 75    ):
 76        names = [k for k in dataset_json if k.startswith(f"{tag}_")]
 77    else:
 78        raise RuntimeError(
 79            "Could not determine the case names for the requested split from 'dataset.json'. The schema of "
 80            "this file could not be verified ahead of time because the dataset repository is gated; please "
 81            "inspect the downloaded 'dataset.json' and adjust '_case_names_for_split' accordingly."
 82        )
 83
 84    names = [n[:-len(".nii.gz")] if n.endswith(".nii.gz") else n for n in names]
 85    return natsorted(set(names))
 86
 87
 88def get_rexgroundingct_metadata(path: Union[os.PathLike, str], download: bool = False) -> Dict[str, Any]:
 89    """Load the per-case finding metadata (free-text descriptions, categories, entity counts) shipped
 90    alongside the ReXGroundingCT 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        The parsed contents of 'dataset.json'.
 98    """
 99    json_path = os.path.join(path, "dataset.json")
100    if not os.path.exists(json_path):
101        if not download:
102            raise RuntimeError(f"Cannot find the data at '{json_path}', but download was set to False.")
103        _download_masks(path, [], download=True, only_json=True)
104
105    with open(json_path, "r") as f:
106        return json.load(f)
107
108
109def _download_masks(path, case_names, download, only_json=False):
110    try:
111        from huggingface_hub import snapshot_download
112    except ImportError:
113        raise ImportError("'huggingface_hub' is required to download ReXGroundingCT. Install it via conda/pip.")
114
115    os.makedirs(path, exist_ok=True)
116    json_path = os.path.join(path, "dataset.json")
117
118    if not os.path.exists(json_path):
119        if not download:
120            raise RuntimeError(f"Cannot find the data at '{json_path}', but download was set to False.")
121        snapshot_download(repo_id=REXGROUNDINGCT_REPO, repo_type="dataset", local_dir=path, allow_patterns="*.json")
122
123    if only_json:
124        return
125
126    missing = [n for n in case_names if not os.path.exists(os.path.join(path, "segmentations", f"{n}.nii.gz"))]
127    if missing:
128        if not download:
129            raise RuntimeError(f"Cannot find {len(missing)} mask(s) at '{path}', but download was set to False.")
130        patterns = [f"segmentations/{n}.nii.gz" for n in missing]
131        snapshot_download(repo_id=REXGROUNDINGCT_REPO, repo_type="dataset", local_dir=path, allow_patterns=patterns)
132
133
134def _download_volumes(path, case_names, download):
135    try:
136        from huggingface_hub import snapshot_download
137    except ImportError:
138        raise ImportError("'huggingface_hub' is required to download CT-RATE. Install it via conda/pip.")
139
140    image_dir = os.path.join(path, "images")
141    os.makedirs(image_dir, exist_ok=True)
142
143    missing, remote_paths = [], {}
144    for name in case_names:
145        remote_path = _ct_rate_path(name)
146        remote_paths[name] = remote_path
147        if not os.path.exists(os.path.join(image_dir, os.path.basename(remote_path))):
148            missing.append(remote_path)
149
150    if missing:
151        if not download:
152            raise RuntimeError(f"Cannot find {len(missing)} CT volume(s) at '{image_dir}', but download was False.")
153        snapshot_download(repo_id=CT_RATE_REPO, repo_type="dataset", local_dir=path, allow_patterns=missing)
154        for remote_path in missing:
155            src = os.path.join(path, remote_path)
156            dst = os.path.join(image_dir, os.path.basename(remote_path))
157            if os.path.exists(src) and not os.path.exists(dst):
158                os.rename(src, dst)
159
160    return image_dir, remote_paths
161
162
163def _merge_instance_mask(mask: np.ndarray) -> np.ndarray:
164    # 'mask' has shape (finding_category, X, Y, Z). Each channel's distinct positive values label
165    # distinct entities of that finding category. Assign a globally unique instance id to every
166    # (category, entity) pair across all channels.
167    merged = np.zeros(mask.shape[1:], dtype="uint16")
168    next_id = 1
169    for channel in mask:
170        for value in np.unique(channel):
171            if value == 0:
172                continue
173            merged[channel == value] = next_id
174            next_id += 1
175    return merged
176
177
178def _merge_masks(mask_dir: str, merged_dir: str, case_names: List[str]) -> None:
179    import nibabel as nib
180
181    os.makedirs(merged_dir, exist_ok=True)
182    for name in case_names:
183        dst = os.path.join(merged_dir, f"{name}.nii.gz")
184        if os.path.exists(dst):
185            continue
186        src = os.path.join(mask_dir, f"{name}.nii.gz")
187        image = nib.load(src)
188        merged = _merge_instance_mask(np.asarray(image.dataobj))
189        nib.save(nib.Nifti1Image(merged, image.affine), dst)
190
191
192def get_rexgroundingct_data(
193    path: Union[os.PathLike, str],
194    split: Literal["train", "val"],
195    max_cases: Optional[int] = None,
196    download: bool = False,
197) -> Tuple[str, str]:
198    """Download the ReXGroundingCT masks and the matching CT-RATE volumes.
199
200    Both HuggingFace repositories are gated (self-service: accept the license on the dataset page while
201    logged in, then pass a HuggingFace access token, e.g. via the `HF_TOKEN` environment variable).
202
203    Args:
204        path: Filepath to a folder where the data is downloaded for further processing.
205        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
206        max_cases: The maximum number of cases to download, taken in order. By default all cases of the
207            requested split are downloaded.
208        download: Whether to download the data if it is not present.
209
210    Returns:
211        Filepath to the folder with the CT volumes.
212        Filepath to the folder with the merged instance segmentation masks.
213    """
214    if split not in SPLITS:
215        raise ValueError(f"'{split}' is not a valid split. Please choose one of {SPLITS}.")
216
217    mask_dir = os.path.join(path, "segmentations")
218    merged_dir = os.path.join(path, "segmentations_merged")
219    dataset_json = get_rexgroundingct_metadata(path, download)
220    case_names = _case_names_for_split(dataset_json, split)
221    if max_cases is not None:
222        case_names = case_names[:max_cases]
223
224    _download_masks(path, case_names, download)
225    image_dir, _ = _download_volumes(path, case_names, download)
226    _merge_masks(mask_dir, merged_dir, case_names)
227
228    return image_dir, merged_dir
229
230
231def get_rexgroundingct_paths(
232    path: Union[os.PathLike, str],
233    split: Literal["train", "val"],
234    max_cases: Optional[int] = None,
235    download: bool = False,
236) -> Tuple[List[str], List[str]]:
237    """Get paths to the ReXGroundingCT data.
238
239    Args:
240        path: Filepath to a folder where the data is downloaded for further processing.
241        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
242        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
243        download: Whether to download the data if it is not present.
244
245    Returns:
246        List of filepaths for the CT volumes.
247        List of filepaths for the finding masks.
248    """
249    image_dir, mask_dir = get_rexgroundingct_data(path, split, max_cases, download)
250
251    label_paths = natsorted(glob(os.path.join(mask_dir, "*.nii.gz")))
252    if max_cases is not None:
253        label_paths = label_paths[:max_cases]
254
255    raw_paths = [os.path.join(image_dir, os.path.basename(p)) for p in label_paths]
256    if len(raw_paths) == 0 or not all(os.path.exists(p) for p in raw_paths):
257        raise RuntimeError("Something went wrong with fetching the image and label paths.")
258
259    return raw_paths, label_paths
260
261
262def get_rexgroundingct_dataset(
263    path: Union[os.PathLike, str],
264    patch_shape: Tuple[int, ...],
265    split: Literal["train", "val"],
266    max_cases: Optional[int] = None,
267    resize_inputs: bool = False,
268    download: bool = False,
269    **kwargs
270) -> Dataset:
271    """Get the ReXGroundingCT dataset for lesion / finding segmentation in chest CT.
272
273    Args:
274        path: Filepath to a folder where the data is downloaded for further processing.
275        patch_shape: The patch shape to use for training.
276        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
277        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
278        resize_inputs: Whether to resize inputs to the desired patch shape.
279        download: Whether to download the data if it is not present.
280        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
281
282    Returns:
283        The segmentation dataset.
284    """
285    raw_paths, label_paths = get_rexgroundingct_paths(path, split, max_cases, download)
286
287    if resize_inputs:
288        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
289        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
290            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
291        )
292
293    return torch_em.default_segmentation_dataset(
294        raw_paths=raw_paths,
295        raw_key="data",
296        label_paths=label_paths,
297        label_key="data",
298        patch_shape=patch_shape,
299        is_seg_dataset=True,
300        **kwargs
301    )
302
303
304def get_rexgroundingct_loader(
305    path: Union[os.PathLike, str],
306    batch_size: int,
307    patch_shape: Tuple[int, ...],
308    split: Literal["train", "val"],
309    max_cases: Optional[int] = None,
310    resize_inputs: bool = False,
311    download: bool = False,
312    **kwargs
313) -> DataLoader:
314    """Get the ReXGroundingCT dataloader for lesion / finding segmentation in chest CT.
315
316    Args:
317        path: Filepath to a folder where the data is downloaded for further processing.
318        batch_size: The batch size for training.
319        patch_shape: The patch shape to use for training.
320        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
321        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
322        resize_inputs: Whether to resize inputs to the desired patch shape.
323        download: Whether to download the data if it is not present.
324        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or the PyTorch DataLoader.
325
326    Returns:
327        The DataLoader.
328    """
329    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
330    dataset = get_rexgroundingct_dataset(path, patch_shape, split, max_cases, resize_inputs, download, **ds_kwargs)
331    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)
REXGROUNDINGCT_REPO = 'rajpurkarlab/ReXGroundingCT'
CT_RATE_REPO = 'ibrahimhamamci/CT-RATE'
SPLITS = ['train', 'val']
SPLIT_TAGS = {'train': 'train', 'val': 'valid'}

Mapping from the split name used by this module to the case name prefix used on the dataset repos.

def get_rexgroundingct_metadata(path: Union[os.PathLike, str], download: bool = False) -> Dict[str, Any]:
 89def get_rexgroundingct_metadata(path: Union[os.PathLike, str], download: bool = False) -> Dict[str, Any]:
 90    """Load the per-case finding metadata (free-text descriptions, categories, entity counts) shipped
 91    alongside the ReXGroundingCT masks.
 92
 93    Args:
 94        path: Filepath to a folder where the data is downloaded for further processing.
 95        download: Whether to download the data if it is not present.
 96
 97    Returns:
 98        The parsed contents of 'dataset.json'.
 99    """
100    json_path = os.path.join(path, "dataset.json")
101    if not os.path.exists(json_path):
102        if not download:
103            raise RuntimeError(f"Cannot find the data at '{json_path}', but download was set to False.")
104        _download_masks(path, [], download=True, only_json=True)
105
106    with open(json_path, "r") as f:
107        return json.load(f)

Load the per-case finding metadata (free-text descriptions, categories, entity counts) shipped alongside the ReXGroundingCT 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:

The parsed contents of 'dataset.json'.

def get_rexgroundingct_data( path: Union[os.PathLike, str], split: Literal['train', 'val'], max_cases: Optional[int] = None, download: bool = False) -> Tuple[str, str]:
193def get_rexgroundingct_data(
194    path: Union[os.PathLike, str],
195    split: Literal["train", "val"],
196    max_cases: Optional[int] = None,
197    download: bool = False,
198) -> Tuple[str, str]:
199    """Download the ReXGroundingCT masks and the matching CT-RATE volumes.
200
201    Both HuggingFace repositories are gated (self-service: accept the license on the dataset page while
202    logged in, then pass a HuggingFace access token, e.g. via the `HF_TOKEN` environment variable).
203
204    Args:
205        path: Filepath to a folder where the data is downloaded for further processing.
206        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
207        max_cases: The maximum number of cases to download, taken in order. By default all cases of the
208            requested split are downloaded.
209        download: Whether to download the data if it is not present.
210
211    Returns:
212        Filepath to the folder with the CT volumes.
213        Filepath to the folder with the merged instance segmentation masks.
214    """
215    if split not in SPLITS:
216        raise ValueError(f"'{split}' is not a valid split. Please choose one of {SPLITS}.")
217
218    mask_dir = os.path.join(path, "segmentations")
219    merged_dir = os.path.join(path, "segmentations_merged")
220    dataset_json = get_rexgroundingct_metadata(path, download)
221    case_names = _case_names_for_split(dataset_json, split)
222    if max_cases is not None:
223        case_names = case_names[:max_cases]
224
225    _download_masks(path, case_names, download)
226    image_dir, _ = _download_volumes(path, case_names, download)
227    _merge_masks(mask_dir, merged_dir, case_names)
228
229    return image_dir, merged_dir

Download the ReXGroundingCT masks and the matching CT-RATE volumes.

Both HuggingFace repositories are gated (self-service: accept the license on the dataset page while logged in, then pass a HuggingFace access token, e.g. via the HF_TOKEN environment variable).

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
  • max_cases: The maximum number of cases to download, taken in order. By default all cases of the requested split are downloaded.
  • download: Whether to download the data if it is not present.
Returns:

Filepath to the folder with the CT volumes. Filepath to the folder with the merged instance segmentation masks.

def get_rexgroundingct_paths( path: Union[os.PathLike, str], split: Literal['train', 'val'], max_cases: Optional[int] = None, download: bool = False) -> Tuple[List[str], List[str]]:
232def get_rexgroundingct_paths(
233    path: Union[os.PathLike, str],
234    split: Literal["train", "val"],
235    max_cases: Optional[int] = None,
236    download: bool = False,
237) -> Tuple[List[str], List[str]]:
238    """Get paths to the ReXGroundingCT data.
239
240    Args:
241        path: Filepath to a folder where the data is downloaded for further processing.
242        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
243        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
244        download: Whether to download the data if it is not present.
245
246    Returns:
247        List of filepaths for the CT volumes.
248        List of filepaths for the finding masks.
249    """
250    image_dir, mask_dir = get_rexgroundingct_data(path, split, max_cases, download)
251
252    label_paths = natsorted(glob(os.path.join(mask_dir, "*.nii.gz")))
253    if max_cases is not None:
254        label_paths = label_paths[:max_cases]
255
256    raw_paths = [os.path.join(image_dir, os.path.basename(p)) for p in label_paths]
257    if len(raw_paths) == 0 or not all(os.path.exists(p) for p in raw_paths):
258        raise RuntimeError("Something went wrong with fetching the image and label paths.")
259
260    return raw_paths, label_paths

Get paths to the ReXGroundingCT data.

Arguments:
  • path: Filepath to a folder where the data is downloaded for further processing.
  • split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
  • max_cases: The maximum number of cases to use. See get_rexgroundingct_data for details.
  • download: Whether to download the data if it is not present.
Returns:

List of filepaths for the CT volumes. List of filepaths for the finding masks.

def get_rexgroundingct_dataset( path: Union[os.PathLike, str], patch_shape: Tuple[int, ...], split: Literal['train', 'val'], max_cases: Optional[int] = None, resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataset.Dataset:
263def get_rexgroundingct_dataset(
264    path: Union[os.PathLike, str],
265    patch_shape: Tuple[int, ...],
266    split: Literal["train", "val"],
267    max_cases: Optional[int] = None,
268    resize_inputs: bool = False,
269    download: bool = False,
270    **kwargs
271) -> Dataset:
272    """Get the ReXGroundingCT dataset for lesion / finding segmentation in chest CT.
273
274    Args:
275        path: Filepath to a folder where the data is downloaded for further processing.
276        patch_shape: The patch shape to use for training.
277        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
278        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
279        resize_inputs: Whether to resize inputs to the desired patch shape.
280        download: Whether to download the data if it is not present.
281        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset`.
282
283    Returns:
284        The segmentation dataset.
285    """
286    raw_paths, label_paths = get_rexgroundingct_paths(path, split, max_cases, download)
287
288    if resize_inputs:
289        resize_kwargs = {"patch_shape": patch_shape, "is_rgb": False}
290        kwargs, patch_shape = util.update_kwargs_for_resize_trafo(
291            kwargs=kwargs, patch_shape=patch_shape, resize_inputs=resize_inputs, resize_kwargs=resize_kwargs
292        )
293
294    return torch_em.default_segmentation_dataset(
295        raw_paths=raw_paths,
296        raw_key="data",
297        label_paths=label_paths,
298        label_key="data",
299        patch_shape=patch_shape,
300        is_seg_dataset=True,
301        **kwargs
302    )

Get the ReXGroundingCT dataset for lesion / finding segmentation in chest CT.

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' (2,992 cases) or 'val' (public validation).
  • max_cases: The maximum number of cases to use. See get_rexgroundingct_data for details.
  • 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_rexgroundingct_loader( path: Union[os.PathLike, str], batch_size: int, patch_shape: Tuple[int, ...], split: Literal['train', 'val'], max_cases: Optional[int] = None, resize_inputs: bool = False, download: bool = False, **kwargs) -> torch.utils.data.dataloader.DataLoader:
305def get_rexgroundingct_loader(
306    path: Union[os.PathLike, str],
307    batch_size: int,
308    patch_shape: Tuple[int, ...],
309    split: Literal["train", "val"],
310    max_cases: Optional[int] = None,
311    resize_inputs: bool = False,
312    download: bool = False,
313    **kwargs
314) -> DataLoader:
315    """Get the ReXGroundingCT dataloader for lesion / finding segmentation in chest CT.
316
317    Args:
318        path: Filepath to a folder where the data is downloaded for further processing.
319        batch_size: The batch size for training.
320        patch_shape: The patch shape to use for training.
321        split: The choice of data split. Either 'train' (2,992 cases) or 'val' (public validation).
322        max_cases: The maximum number of cases to use. See `get_rexgroundingct_data` for details.
323        resize_inputs: Whether to resize inputs to the desired patch shape.
324        download: Whether to download the data if it is not present.
325        kwargs: Additional keyword arguments for `torch_em.default_segmentation_dataset` or the PyTorch DataLoader.
326
327    Returns:
328        The DataLoader.
329    """
330    ds_kwargs, loader_kwargs = util.split_kwargs(torch_em.default_segmentation_dataset, **kwargs)
331    dataset = get_rexgroundingct_dataset(path, patch_shape, split, max_cases, resize_inputs, download, **ds_kwargs)
332    return torch_em.get_data_loader(dataset, batch_size, **loader_kwargs)

Get the ReXGroundingCT dataloader for lesion / finding segmentation in chest CT.

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' (2,992 cases) or 'val' (public validation).
  • max_cases: The maximum number of cases to use. See get_rexgroundingct_data for details.
  • 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 the PyTorch DataLoader.
Returns:

The DataLoader.