torch_em.data.datasets.util

  1import os
  2import hashlib
  3import inspect
  4import zipfile
  5import requests
  6from tqdm import tqdm
  7from warnings import warn
  8from subprocess import run
  9from xml.dom import minidom
 10from packaging import version
 11from shutil import copyfileobj, which
 12
 13from typing import Optional, Tuple, Literal, List, Dict, Union, Callable
 14
 15import numpy as np
 16from skimage.draw import polygon
 17
 18import torch
 19
 20import torch_em
 21from torch_em.transform import get_raw_transform
 22from torch_em.transform.generic import ResizeLongestSideInputs, Compose
 23
 24try:
 25    import gdown
 26except ImportError:
 27    gdown = None
 28
 29try:
 30    from tcia_utils import nbia
 31except ModuleNotFoundError:
 32    nbia = None
 33
 34try:
 35    from cryoet_data_portal import Client, Dataset
 36except ImportError:
 37    Client, Dataset = None, None
 38
 39try:
 40    import synapseclient
 41    import synapseutils
 42except ImportError:
 43    synapseclient, synapseutils = None, None
 44
 45
 46BIOIMAGEIO_IDS = {
 47    "covid_if": "ilastik/covid_if_training_data",
 48    "cremi": "ilastik/cremi_training_data",
 49    "dsb": "ilastik/stardist_dsb_training_data",
 50    "hpa": "",  # not on bioimageio yet
 51    "isbi2012": "ilastik/isbi2012_neuron_segmentation_challenge",
 52    "kasthuri": "",  # not on bioimageio yet:
 53    "livecell": "ilastik/livecell_dataset",
 54    "lucchi": "",  # not on bioimageio yet:
 55    "mitoem": "ilastik/mitoem_segmentation_challenge",
 56    "monuseg": "deepimagej/monuseg_digital_pathology_miccai2018",
 57    "ovules": "",  # not on bioimageio yet
 58    "plantseg_root": "ilastik/plantseg_root",
 59    "plantseg_ovules": "ilastik/plantseg_ovules",
 60    "platynereis": "ilastik/platynereis_em_training_data",
 61    "snemi": "",  # not on bioimagegio yet
 62    "uro_cell": "",  # not on bioimageio yet: https://doi.org/10.1016/j.compbiomed.2020.103693
 63    "vnc": "ilastik/vnc",
 64}
 65"""@private
 66"""
 67
 68
 69def get_bioimageio_dataset_id(dataset_name):
 70    """@private
 71    """
 72    assert dataset_name in BIOIMAGEIO_IDS
 73    return BIOIMAGEIO_IDS[dataset_name]
 74
 75
 76def get_checksum(filename: str) -> str:
 77    """Get the SHA256 checksum of a file.
 78
 79    Args:
 80        filename: The filepath.
 81
 82    Returns:
 83        The checksum.
 84    """
 85    # The file is hashed in chunks, so that datasets with multi-GB archives do not run out of memory.
 86    hasher = hashlib.sha256()
 87    with open(filename, "rb") as f:
 88        for chunk in iter(lambda: f.read(64 * 1024 * 1024), b""):
 89            hasher.update(chunk)
 90    return hasher.hexdigest()
 91
 92
 93def _check_checksum(path, checksum):
 94    if checksum is not None:
 95        this_checksum = get_checksum(path)
 96        if this_checksum != checksum:
 97            raise RuntimeError(
 98                "The checksum of the download does not match the expected checksum."
 99                f"Expected: {checksum}, got: {this_checksum}"
100            )
101        print("Download successful and checksums agree.")
102    else:
103        warn("The file was downloaded, but no checksum was provided, so the file may be corrupted.")
104
105
106# this needs to be extended to support download from s3 via boto,
107# if we get a resource that is available via s3 without support for http
108def download_source(path: str, url: str, download: bool, checksum: Optional[str] = None, verify: bool = True) -> None:
109    """Download data via https.
110
111    Args:
112        path: The path for saving the data.
113        url: The url of the data.
114        download: Whether to download the data if it is not saved at `path` yet.
115        checksum: The expected checksum of the data.
116        verify: Whether to verify the https address.
117    """
118    if os.path.exists(path):
119        return
120    if not download:
121        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
122
123    # The data is downloaded to a temporary path and only moved to `path` once it is complete and verified.
124    # Otherwise an interrupted download would be mistaken for a complete one by the check above.
125    tmp_path = f"{path}.incomplete"
126    with requests.get(url, stream=True, allow_redirects=True, verify=verify) as r:
127        r.raise_for_status()  # check for error
128        # Compute checksums on the file content rather than its HTTP transfer encoding.
129        r.raw.decode_content = True
130        file_size = int(r.headers.get("Content-Length", 0))
131        desc = f"Download {url} to {path}"
132        if file_size == 0:
133            desc += " (unknown file size)"
134        with tqdm.wrapattr(r.raw, "read", total=file_size, desc=desc) as r_raw, open(tmp_path, "wb") as f:
135            copyfileobj(r_raw, f)
136
137    _check_checksum(tmp_path, checksum)
138    os.replace(tmp_path, path)
139
140
141def download_source_gdrive(
142    path: str,
143    url: str,
144    download: bool,
145    checksum: Optional[str] = None,
146    download_type: Literal["zip", "folder"] = "zip",
147    expected_samples: int = 10000,
148    quiet: bool = True,
149) -> None:
150    """Download data from google drive.
151
152    Args:
153        path: The path for saving the data.
154        url: The url of the data.
155        download: Whether to download the data if it is not saved at `path` yet.
156        checksum: The expected checksum of the data.
157        download_type: The download type, either 'zip' or 'folder'.
158        expected_samples: The maximal number of samples in the folder.
159        quiet: Whether to download quietly.
160    """
161    if os.path.exists(path):
162        return
163
164    if not download:
165        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
166
167    if gdown is None:
168        raise RuntimeError(
169            "Need gdown library to download data from google drive. "
170            "Please install gdown: 'conda install -c conda-forge gdown==4.6.3'."
171        )
172
173    print("Downloading the files. Might take a few minutes...")
174
175    if download_type == "zip":
176        gdown.download(url, path, quiet=quiet)
177        _check_checksum(path, checksum)
178    elif download_type == "folder":
179        assert version.parse(gdown.__version__) == version.parse("4.6.3"), "Please install 'gdown==4.6.3'."
180        gdown.download_folder.__globals__["MAX_NUMBER_FILES"] = expected_samples
181        gdown.download_folder(url=url, output=path, quiet=quiet, remaining_ok=True)
182    else:
183        raise ValueError("`download_path` argument expects either `zip`/`folder`")
184
185    print("Download completed.")
186
187
188def download_source_empiar(path: str, access_id: str, download: bool) -> str:
189    """Download data from EMPIAR.
190
191    Requires the ascp command from the aspera CLI.
192
193    Args:
194        path: The path for saving the data.
195        access_id: The EMPIAR accession id of the data to download.
196        download: Whether to download the data if it is not saved at `path` yet.
197
198    Returns:
199        The path to the downloaded data.
200    """
201    download_path = os.path.join(path, access_id)
202
203    if os.path.exists(download_path):
204        return download_path
205    if not download:
206        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
207
208    if which("ascp") is None:
209        raise RuntimeError(
210            "Need aspera-cli to download data from empiar. You can install it via 'conda install -c hcc aspera-cli'."
211        )
212
213    key_file = os.path.expanduser("~/.aspera/cli/etc/asperaweb_id_dsa.openssh")
214    if not os.path.exists(key_file):
215        conda_root = os.environ["CONDA_PREFIX"]
216        key_file = os.path.join(conda_root, "etc/asperaweb_id_dsa.openssh")
217
218    if not os.path.exists(key_file):
219        raise RuntimeError("Could not find the aspera ssh keyfile")
220
221    cmd = ["ascp", "-QT", "-l", "200M", "-P33001", "-i", key_file, f"emp_ext2@fasp.ebi.ac.uk:/{access_id}", path]
222    run(cmd)
223
224    return download_path
225
226
227def download_source_kaggle(path: str, dataset_name: str, download: bool, competition: bool = False):
228    """Download data from Kaggle.
229
230    Requires the Kaggle API.
231
232    Args:
233        path: The path for saving the data.
234        dataset_name: The name of the dataset to download.
235        download: Whether to download the data if it is not saved at `path` yet.
236        competition: Whether this data is from a competition and requires the kaggle.competition API.
237    """
238    if not download:
239        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
240
241    try:
242        from kaggle.api.kaggle_api_extended import KaggleApi
243    except ModuleNotFoundError:
244        msg = "Please install the Kaggle API. You can do this using 'pip install kaggle'. "
245        msg += "After you have installed kaggle, you would need an API token. "
246        msg += "Follow the instructions at https://www.kaggle.com/docs/api."
247        raise ModuleNotFoundError(msg)
248
249    api = KaggleApi()
250    api.authenticate()
251
252    if competition:
253        api.competition_download_files(competition=dataset_name, path=path, quiet=False)
254    else:
255        api.dataset_download_files(dataset=dataset_name, path=path, quiet=False)
256
257
258NBIA_API_URL = "https://services.cancerimagingarchive.net/nbia-api/services/v1/"
259
260
261def _download_tcia_series_with_rest(series_uids, dst, csv_filename):
262    """Download DICOM series from TCIA via the NBIA REST API.
263
264    This is the fallback for missing 'tcia_utils'. It mimics the on-disk layout of 'nbia.downloadSeries':
265    each series is extracted to '<dst>/<SeriesInstanceUID>/' and the series metadata are written to
266    '<csv_filename>.csv' (including the 'Series UID', 'Subject ID' and 'Modality' columns).
267    The metadata of each series is cached in '<csv_filename>.partial.json' while the download is running,
268    so that an interrupted download does not have to query the metadata of all series again.
269    See https://wiki.cancerimagingarchive.net/x/fILTB for the API documentation.
270    """
271    import csv
272    import json
273    import time
274    import tempfile
275
276    def get_with_retries(endpoint, n_retries=5, **kwargs):
277        # The NBIA API occasionally returns server errors, so the requests are retried.
278        for attempt in range(n_retries):
279            response = requests.get(NBIA_API_URL + endpoint, **kwargs)
280            if response.status_code < 500 or attempt == n_retries - 1:
281                response.raise_for_status()
282                return response
283            response.close()
284            time.sleep(10 * (attempt + 1))
285
286    os.makedirs(dst, exist_ok=True)
287    cache_path = f"{csv_filename}.partial.json"
288    cache = {}
289    if os.path.exists(cache_path):
290        with open(cache_path, "r") as f:
291            cache = json.load(f)
292
293    metadata = []
294    for i, uid in enumerate(tqdm(series_uids, desc=f"Download {len(series_uids)} series from TCIA to {dst}")):
295        series_dir = os.path.join(dst, uid)
296        if uid in cache and os.path.exists(series_dir):
297            metadata.append(cache[uid])
298            continue
299
300        response = get_with_retries("getSeriesMetaData", params={"SeriesInstanceUID": uid})
301        # A small number of series have no indexed metadata and the endpoint returns an empty body for
302        # them, even though the image data itself downloads without issue. A transient proxy or gateway
303        # error can also return a non-JSON body with a 200 status, which is handled the same way.
304        try:
305            rows = response.json() if response.content else []
306        except ValueError:
307            rows = []
308        metadata.append(rows[0] if rows else {"Series UID": uid})
309        cache[uid] = metadata[-1]
310        if i % 50 == 0:
311            with open(cache_path, "w") as f:
312                json.dump(cache, f)
313
314        if os.path.exists(series_dir):  # This series has been downloaded already.
315            continue
316
317        # The series is downloaded as a zip archive, which is extracted to a temporary folder
318        # and only moved to the final location once it is complete.
319        try:
320            with tempfile.TemporaryDirectory(dir=dst) as tmp_dir:
321                zip_path = os.path.join(tmp_dir, "series.zip")
322                with get_with_retries("getImage", params={"SeriesInstanceUID": uid}, stream=True) as r:
323                    with open(zip_path, "wb") as f:
324                        copyfileobj(r.raw, f)
325                tmp_series_dir = os.path.join(tmp_dir, "series")
326                unzip(zip_path, tmp_series_dir)
327                os.rename(tmp_series_dir, series_dir)
328        except FileExistsError:
329            # A resumed download may have already extracted this series: on some network filesystems
330            # the 'os.path.exists' check above can be stale, so this is not caught earlier.
331            pass
332        except requests.exceptions.HTTPError as e:
333            # A small number of series are indexed but no longer resolvable through this endpoint (e.g. a
334            # stale UID), which should not abort downloading the rest of a possibly multi-hour batch.
335            if e.response is not None and e.response.status_code < 500:
336                tqdm.write(f"Skipping {uid}, which failed to download: {e}")
337            else:
338                raise
339
340    # The metadata keys differ between series (e.g. 'Series Date' is only reported for some), so the header
341    # has to be the union of all keys.
342    fieldnames = list(dict.fromkeys(key for row in metadata for key in row))
343    csv_path = f"{csv_filename}.csv"
344    with open(csv_path, "w", newline="") as f:
345        writer = csv.DictWriter(f, fieldnames=fieldnames)
346        writer.writeheader()
347        writer.writerows(metadata)
348
349    if os.path.exists(cache_path):  # The download is complete, so the metadata cache is not needed anymore.
350        os.remove(cache_path)
351
352
353def _download_tcia_manifest_with_rest(manifest_path, dst, csv_filename):
354    """Download all series listed in a TCIA manifest via the NBIA REST API."""
355    with open(manifest_path, "r") as f:
356        lines = [line.strip() for line in f.readlines()]
357    series_uids = lines[lines.index("ListOfSeriesToDownload=") + 1:]
358    series_uids = [uid for uid in series_uids if uid]
359    _download_tcia_series_with_rest(series_uids, dst, csv_filename)
360
361
362def download_tcia_series(series_uids: List[str], dst: str, csv_filename: str) -> str:
363    """Download individual DICOM series from TCIA by their series instance UIDs.
364
365    Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise.
366    Each series is stored in '<dst>/<SeriesInstanceUID>/', series that exist there already are skipped.
367
368    Args:
369        series_uids: The UIDs of the series to download.
370        dst: The folder for saving the DICOM series.
371        csv_filename: The path for saving the series metadata (without the '.csv' extension).
372
373    Returns:
374        The path to the csv file with the series metadata.
375    """
376    if nbia is None:
377        _download_tcia_series_with_rest(series_uids, dst, csv_filename)
378    else:
379        nbia.downloadSeries(series_data=series_uids, input_type="list", path=dst, csv_filename=csv_filename)
380    return f"{csv_filename}.csv"
381
382
383def download_source_tcia(path, url, dst, csv_filename, download):
384    """Download data from TCIA.
385
386    Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise.
387
388    Args:
389        path: The path for saving the manifest file. If `url` is None, this must point to an existing manifest,
390            e.g. one that was written by the caller to download only a subset of the series of a collection.
391        url: The URL to the TCIA manifest of the dataset. Set to None to use the manifest at `path`.
392        dst: The folder for saving the DICOM series. Each series is stored in a sub-folder named after its UID.
393        csv_filename: The path for saving the series metadata (without the '.csv' extension).
394        download: Whether to download the data if it is not saved at `path` yet.
395    """
396    if not download:
397        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
398
399    if url is None:
400        assert os.path.exists(path), f"The manifest {path} does not exist."
401    else:
402        assert url.endswith(".tcia"), f"{url} is not a TCIA Manifest."
403        # Downloads the manifest file from the collection page.
404        manifest = requests.get(url=url)
405        manifest.raise_for_status()
406        with open(path, "wb") as f:
407            f.write(manifest.content)
408
409    # This part extracts the UIDs from the manifests and downloads them.
410    if nbia is None:
411        _download_tcia_manifest_with_rest(path, dst, csv_filename)
412    else:
413        nbia.downloadSeries(series_data=path, input_type="manifest", path=dst, csv_filename=csv_filename)
414
415
416def download_source_synapse(path: str, entity: str, download: bool) -> None:
417    """Download data from synapse.
418
419    Requires the synapseclient python library.
420
421    Args:
422        path: The path for saving the data.
423        entity: The name of the data to download from synapse.
424        download: Whether to download the data if it is not saved at `path` yet.
425    """
426    if not download:
427        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
428
429    if synapseclient is None:
430        raise RuntimeError(
431            "You must install 'synapseclient' to download files from 'synapse'. "
432            "Remember to create an account and generate an authentication code for your account. "
433            "Please follow the documentation for details on creating the '~/.synapseConfig' file here: "
434            "https://python-docs.synapse.org/tutorials/authentication/."
435        )
436
437    assert entity.startswith("syn"), "The entity name does not look as expected. It should be something like 'syn123'."
438
439    # Download all files in the folder.
440    syn = synapseclient.Synapse()
441    syn.login()  # Since we do not pass any credentials here, it fetches all details from '~/.synapseConfig'.
442    synapseutils.syncFromSynapse(syn=syn, entity=entity, path=path)
443
444
445def update_kwargs(kwargs, key, value, msg=None):
446    """@private
447    """
448    if key in kwargs:
449        msg = f"{key} will be over-ridden in loader kwargs." if msg is None else msg
450        warn(msg)
451    kwargs[key] = value
452    return kwargs
453
454
455def unzip_tarfile(tar_path: str, dst: str, remove: bool = True) -> None:
456    """Unpack a tar archive.
457
458    Args:
459        tar_path: Path to the tar file.
460        dst: Where to unpack the archive.
461        remove: Whether to remove the tar file after unpacking.
462    """
463    import tarfile
464
465    if tar_path.endswith(".tar.gz") or tar_path.endswith(".tgz"):
466        access_mode = "r:gz"
467    elif tar_path.endswith(".tar"):
468        access_mode = "r:"
469    else:
470        raise ValueError(f"The provided file isn't a supported archive to unpack. Please check the file: {tar_path}.")
471
472    tar = tarfile.open(tar_path, access_mode)
473    tar.extractall(dst)
474    tar.close()
475
476    if remove:
477        os.remove(tar_path)
478
479
480def unzip_rarfile(rar_path: str, dst: str, remove: bool = True, use_rarfile: bool = True) -> None:
481    """Unpack a rar archive.
482
483    Args:
484        rar_path: Path to the rar file.
485        dst: Where to unpack the archive.
486        remove: Whether to remove the tar file after unpacking.
487        use_rarfile: Whether to use the rarfile library or aspose.zip.
488    """
489    def _extract_with_rarfile():
490        import rarfile
491        with rarfile.RarFile(rar_path) as archive:
492            archive.extractall(path=dst)
493
494    def _extract_with_aspose():
495        import aspose.zip as az
496        with az.rar.RarArchive(rar_path) as archive:
497            archive.extract_to_directory(dst)
498
499    def _extract_with_7z():
500        if which("7z") is None:
501            raise RuntimeError("The 'p7zip' CLI is not available.")
502        run(["7z", "x", f"-o{dst}", "-y", rar_path], check=True)
503
504    extractors = [
505        ('rarfile', _extract_with_rarfile), ('aspose.zip', _extract_with_aspose), ('7z', _extract_with_7z),
506    ] if use_rarfile else [('aspose.zip', _extract_with_aspose), ('7z', _extract_with_7z)]
507
508    errors = []
509    for name, extractor in extractors:
510        try:
511            extractor()
512            break
513        except Exception as err:
514            errors.append((name, err))
515            if len(errors) < len(extractors):
516                next_name = extractors[len(errors)][0]
517                warn(f"Extraction with '{name}' failed for {rar_path} ({err}). Falling back to '{next_name}'.")
518    else:
519        backends = ', '.join(f"'{name}'" for name, _ in extractors)
520        raise RuntimeError(
521            f"Failed to extract rar archive {rar_path} with {backends}. "
522            "Please ensure one of the supported backends is installed and can read this archive."
523        ) from errors[-1][1]
524
525    if remove:
526        os.remove(rar_path)
527
528
529def unzip(zip_path: str, dst: str, remove: bool = True) -> None:
530    """Unpack a zip archive.
531
532    Args:
533        zip_path: Path to the zip file.
534        dst: Where to unpack the archive.
535        remove: Whether to remove the tar file after unpacking.
536    """
537    with zipfile.ZipFile(zip_path, "r") as f:
538        f.extractall(dst)
539    if remove:
540        os.remove(zip_path)
541
542
543def unzip_7z(path_7z: str, dst: str, remove: bool = True) -> None:
544    """Unpack a 7z archive.
545
546    Args:
547        path_7z: Path to the 7z file.
548        dst: Where to unpack the archive.
549        remove: Whether to remove the 7z file after unpacking.
550    """
551    if which("7z") is None:
552        raise RuntimeError("Need the 'p7zip' CLI to extract 7z archives. You can install it via 'conda install -c conda-forge p7zip'.")  # noqa
553
554    run(["7z", "x", f"-o{dst}", "-y", path_7z])
555
556    if remove:
557        os.remove(path_7z)
558
559
560def split_kwargs(function, **kwargs):
561    """@private
562    """
563    function_parameters = inspect.signature(function).parameters
564    parameter_names = list(function_parameters.keys())
565    other_kwargs = {k: v for k, v in kwargs.items() if k not in parameter_names}
566    kwargs = {k: v for k, v in kwargs.items() if k in parameter_names}
567    return kwargs, other_kwargs
568
569
570# this adds the default transforms for 'raw_transform' and 'transform'
571# in case these were not specified in the kwargs
572# this is NOT necessary if 'default_segmentation_dataset' is used, only if a dataset class
573# is used directly, e.g. in the LiveCell Loader
574def ensure_transforms(ndim, **kwargs):
575    """@private
576    """
577    if "raw_transform" not in kwargs:
578        kwargs = update_kwargs(kwargs, "raw_transform", torch_em.transform.get_raw_transform())
579    if "transform" not in kwargs:
580        kwargs = update_kwargs(kwargs, "transform", torch_em.transform.get_augmentations(ndim=ndim))
581    return kwargs
582
583
584def add_instance_label_transform(
585    kwargs, add_binary_target, label_dtype=None, binary=False, boundaries=False, offsets=None, binary_is_exclusive=True,
586):
587    """@private
588    """
589    if binary_is_exclusive:
590        assert sum((offsets is not None, boundaries, binary)) <= 1
591    else:
592        assert sum((offsets is not None, boundaries)) <= 1
593    if offsets is not None:
594        label_transform2 = torch_em.transform.label.AffinityTransform(offsets=offsets,
595                                                                      add_binary_target=add_binary_target,
596                                                                      add_mask=True)
597        msg = "Offsets are passed, but 'label_transform2' is in the kwargs. It will be over-ridden."
598        kwargs = update_kwargs(kwargs, "label_transform2", label_transform2, msg=msg)
599        label_dtype = torch.float32
600    elif boundaries:
601        label_transform = torch_em.transform.label.BoundaryTransform(add_binary_target=add_binary_target)
602        msg = "Boundaries is set to true, but 'label_transform' is in the kwargs. It will be over-ridden."
603        kwargs = update_kwargs(kwargs, "label_transform", label_transform, msg=msg)
604        label_dtype = torch.float32
605    elif binary:
606        label_transform = torch_em.transform.label.labels_to_binary
607        msg = "Binary is set to true, but 'label_transform' is in the kwargs. It will be over-ridden."
608        kwargs = update_kwargs(kwargs, "label_transform", label_transform, msg=msg)
609        label_dtype = torch.float32
610    return kwargs, label_dtype
611
612
613def update_kwargs_for_resize_trafo(kwargs, patch_shape, resize_inputs, resize_kwargs=None, ensure_rgb=None):
614    """@private
615    """
616    # Checks for raw_transform and label_transform incoming values.
617    # If yes, it will automatically merge these two transforms to apply them together.
618    if resize_inputs:
619        assert isinstance(resize_kwargs, dict)
620
621        target_shape = resize_kwargs.get("patch_shape")
622        if len(resize_kwargs["patch_shape"]) == 3:
623            # we only need the XY dimensions to reshape the inputs along them.
624            target_shape = target_shape[1:]
625            # we provide the Z dimension value to return the desired number of slices and not the whole volume
626            kwargs["z_ext"] = resize_kwargs["patch_shape"][0]
627
628        raw_trafo = ResizeLongestSideInputs(target_shape=target_shape, is_rgb=resize_kwargs["is_rgb"])
629        label_trafo = ResizeLongestSideInputs(target_shape=target_shape, is_label=True)
630
631        # The patch shape provided to the dataset. Here, "None" means that the entire volume will be loaded.
632        patch_shape = None
633
634    if ensure_rgb is None:
635        raw_trafos = []
636    else:
637        assert not isinstance(ensure_rgb, bool), "'ensure_rgb' is expected to be a function."
638        raw_trafos = [ensure_rgb]
639
640    if "raw_transform" in kwargs:
641        raw_trafos.extend([raw_trafo, kwargs["raw_transform"]])
642    else:
643        raw_trafos.extend([raw_trafo, get_raw_transform()])
644
645    kwargs["raw_transform"] = Compose(*raw_trafos, is_multi_tensor=False)
646
647    if "label_transform" in kwargs:
648        trafo = Compose(label_trafo, kwargs["label_transform"], is_multi_tensor=False)
649        kwargs["label_transform"] = trafo
650    else:
651        kwargs["label_transform"] = label_trafo
652
653    return kwargs, patch_shape
654
655
656def generate_labeled_array_from_xml(shape: Tuple[int, ...], xml_file: str) -> np.ndarray:
657    """Generate a label mask from a contour defined in a xml annotation file.
658
659    Function taken from: https://github.com/rshwndsz/hover-net/blob/master/lightning_hovernet.ipynb
660
661    Args:
662        shape: The image shape.
663        xml_file: The path to the xml file with contour annotations.
664
665    Returns:
666        The label mask.
667    """
668    # DOM object created by the minidom parser
669    xDoc = minidom.parse(xml_file)
670
671    # List of all Region tags
672    regions = xDoc.getElementsByTagName('Region')
673
674    # List which will store the vertices for each region
675    xy = []
676    for region in regions:
677        # Loading all the vertices in the region
678        vertices = region.getElementsByTagName('Vertex')
679
680        # The vertices of a region will be stored in a array
681        vw = np.zeros((len(vertices), 2))
682
683        for index, vertex in enumerate(vertices):
684            # Storing the values of x and y coordinate after conversion
685            vw[index][0] = float(vertex.getAttribute('X'))
686            vw[index][1] = float(vertex.getAttribute('Y'))
687
688        # Append the vertices of a region
689        xy.append(np.int32(vw))
690
691    # Creating a completely black image
692    mask = np.zeros(shape, np.uint32)  # Integer instance ids; float labels break connected-component ops.
693
694    # Start the instance ids at 1: id 0 is background, so enumerating from 0 silently drops the first region.
695    for i, contour in enumerate(xy, start=1):
696        r, c = polygon(np.array(contour)[:, 1], np.array(contour)[:, 0], shape=shape)
697        mask[r, c] = i
698    return mask
699
700
701def load_dicom_series(series_dir: str) -> Tuple[np.ndarray, Dict[str, np.ndarray]]:
702    """Stack a single-frame DICOM image series (CT, MR, PET) into a volume with axes (z, y, x).
703
704    The slices are sorted by their position along the slice normal (the cross product of the row and column
705    direction in 'ImageOrientationPatient'), so the volume is stacked consistently for any acquisition plane.
706    'RescaleSlope' and 'RescaleIntercept' are applied per slice, i.e. CT volumes are returned in Hounsfield units.
707
708    NOTE: This requires the pydicom python package.
709
710    Args:
711        series_dir: The folder with the DICOM files of the series.
712
713    Returns:
714        The volume with axes (z, y, x) as float32.
715        The geometry of the volume, which is needed by `rasterize_rtstruct`. A dictionary with the keys
716        'origin' (the 'ImagePositionPatient' of each slice, n_slices x 3), 'row_direction' and 'column_direction'
717        (the unit vectors along which the column and the row index increase, from 'ImageOrientationPatient'),
718        'spacing' (the row and column spacing from 'PixelSpacing') and 'sop_uids' (the 'SOPInstanceUID' per slice).
719    """
720    import pydicom
721
722    dcm_paths = [os.path.join(series_dir, fname) for fname in sorted(os.listdir(series_dir)) if fname.endswith(".dcm")]
723    slices = [pydicom.dcmread(dcm_path) for dcm_path in dcm_paths]
724
725    row_direction = np.array([float(v) for v in slices[0].ImageOrientationPatient[:3]])
726    column_direction = np.array([float(v) for v in slices[0].ImageOrientationPatient[3:]])
727    normal = np.cross(row_direction, column_direction)
728    slices.sort(key=lambda dcm: np.dot([float(v) for v in dcm.ImagePositionPatient], normal))
729
730    volume = []
731    for dcm in slices:
732        frame = dcm.pixel_array.astype("float32")
733        volume.append(frame * float(dcm.get("RescaleSlope", 1.0)) + float(dcm.get("RescaleIntercept", 0.0)))
734    volume = np.stack(volume)
735
736    geometry = {
737        "origin": np.array([[float(v) for v in dcm.ImagePositionPatient] for dcm in slices]),
738        "row_direction": row_direction,
739        "column_direction": column_direction,
740        "spacing": np.array([float(v) for v in slices[0].PixelSpacing]),
741        "sop_uids": np.array([str(dcm.SOPInstanceUID) for dcm in slices]),
742    }
743    return volume, geometry
744
745
746def rasterize_rtstruct(
747    rtstruct_path: str,
748    geometry: Dict[str, np.ndarray],
749    shape: Tuple[int, int, int],
750    roi_labels: Union[Dict[str, int], Callable[[int, str], Optional[int]]],
751) -> np.ndarray:
752    """Rasterize the contours of a DICOM RTSTRUCT file onto the voxel grid of the referenced image series.
753
754    Each 'CLOSED_PLANAR' contour is assigned to the slice it references (via 'ReferencedSOPInstanceUID', with a
755    fallback to the closest slice along the slice normal). Its points are projected onto the row and column
756    direction of that slice to obtain pixel coordinates and the polygon is filled with `skimage.draw.polygon`.
757    Multiple contours of the same ROI on the same slice are combined with XOR, so that inner contours form holes.
758    Where different ROIs overlap, the ROI with the lower label id takes precedence.
759
760    NOTE: This requires the pydicom python package.
761
762    Args:
763        rtstruct_path: The path to the RTSTRUCT DICOM file.
764        geometry: The geometry of the referenced image series, as returned by `load_dicom_series`.
765        shape: The shape of the image volume (z, y, x).
766        roi_labels: The mapping from ROIs to label ids. Either a dictionary that maps the ROI names to label ids,
767            or a function that maps the ROI number and ROI name to a label id. ROIs that are not in the dictionary
768            or for which the function returns None are ignored.
769
770    Returns:
771        The label volume (uint8) with axes (z, y, x).
772    """
773    import pydicom
774
775    rtstruct = pydicom.dcmread(rtstruct_path)
776    roi_names = {int(roi.ROINumber): str(roi.ROIName) for roi in rtstruct.StructureSetROISequence}
777    slice_ids = {uid: z for z, uid in enumerate(geometry["sop_uids"])}
778    normal = np.cross(geometry["row_direction"], geometry["column_direction"])
779    slice_positions = geometry["origin"] @ normal
780
781    masks = {}
782    for roi_contour in rtstruct.ROIContourSequence:
783        roi_number = int(roi_contour.ReferencedROINumber)
784        roi_name = roi_names[roi_number]
785        if isinstance(roi_labels, dict):
786            label_id = roi_labels.get(roi_name)
787        else:
788            label_id = roi_labels(roi_number, roi_name)
789        if label_id is None:
790            continue
791
792        mask = masks.setdefault(label_id, np.zeros(shape, dtype="bool"))
793        for contour in roi_contour.get("ContourSequence", []):
794            if contour.ContourGeometricType != "CLOSED_PLANAR":
795                continue
796            points = np.array([float(v) for v in contour.ContourData]).reshape(-1, 3)
797            if len(points) < 3:
798                continue
799
800            z = None
801            for ref in contour.get("ContourImageSequence", []):
802                z = slice_ids.get(str(ref.ReferencedSOPInstanceUID))
803            if z is None:
804                z = int(np.argmin(np.abs(slice_positions - np.dot(points[0], normal))))
805
806            offsets = points - geometry["origin"][z]
807            cols = offsets @ geometry["row_direction"] / geometry["spacing"][1]
808            rows = offsets @ geometry["column_direction"] / geometry["spacing"][0]
809            rr, cc = polygon(rows, cols, shape=shape[1:])
810            mask[z, rr, cc] = ~mask[z, rr, cc]
811
812    labels = np.zeros(shape, dtype="uint8")
813    for label_id in sorted(masks, reverse=True):
814        labels[masks[label_id]] = label_id
815    return labels
816
817
818# This function could be extended to convert WSIs (or modalities with multiple resolutions).
819def convert_svs_to_array(
820    path: str, location: Tuple[int, int] = (0, 0), level: int = 0, img_size: Tuple[int, int] = None,
821) -> np.ndarray:
822    """Convert a .svs file for WSI imagging to a numpy array.
823
824    Requires the tiffslide python library.
825    The function can load multi-resolution images. You can specify the resolution level via `level`.
826
827    Args:
828        path: File path ath to the svs file.
829        location: Pixel location (x, y) in level 0 of the image.
830        level: Target level used to read the image.
831        img_size: Size of the image. If None, the shape of the image at `level` is used.
832
833    Returns:
834        The image as numpy array.
835    """
836    assert path.endswith(".svs"), f"The provided file ({path}) isn't in svs format"
837
838    try:
839        from tiffslide import TiffSlide
840    except ImportError:
841        # svs is a pyramidal TIFF variant, so tifffile can read it without the tiffslide dependency.
842        import tifffile
843        with tifffile.TiffFile(path) as f:
844            image = f.series[0].levels[level].asarray()
845        x, y = location
846        if img_size is not None:
847            image = image[y:y + img_size[1], x:x + img_size[0]]
848        else:
849            image = image[y:, x:]
850        return image
851
852    _slide = TiffSlide(path)
853    if img_size is None:
854        img_size = _slide.level_dimensions[0]
855    return _slide.read_region(location=location, level=level, size=img_size, as_array=True)
856
857
858def download_from_cryo_et_portal(path: str, dataset_id: int, download: bool) -> str:
859    """Download data from the CryoET Data Portal.
860
861    Requires the cryoet-data-portal python library.
862
863    Args:
864        path: The path for saving the data.
865        dataset_id: The id of the data to download from the portal.
866        download: Whether to download the data if it is not saved at `path` yet.
867
868    Returns:
869        The file path to the downloaded data.
870    """
871    if Client is None or Dataset is None:
872        raise RuntimeError("Please install CryoETDataPortal via 'pip install cryoet-data-portal'")
873
874    output_path = os.path.join(path, str(dataset_id))
875    if os.path.exists(output_path):
876        return output_path
877
878    if not download:
879        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
880
881    client = Client()
882    dataset = Dataset.get_by_id(client, dataset_id)
883    dataset.download_everything(dest_path=path)
884
885    return output_path
def get_checksum(filename: str) -> str:
77def get_checksum(filename: str) -> str:
78    """Get the SHA256 checksum of a file.
79
80    Args:
81        filename: The filepath.
82
83    Returns:
84        The checksum.
85    """
86    # The file is hashed in chunks, so that datasets with multi-GB archives do not run out of memory.
87    hasher = hashlib.sha256()
88    with open(filename, "rb") as f:
89        for chunk in iter(lambda: f.read(64 * 1024 * 1024), b""):
90            hasher.update(chunk)
91    return hasher.hexdigest()

Get the SHA256 checksum of a file.

Arguments:
  • filename: The filepath.
Returns:

The checksum.

def download_source( path: str, url: str, download: bool, checksum: Optional[str] = None, verify: bool = True) -> None:
109def download_source(path: str, url: str, download: bool, checksum: Optional[str] = None, verify: bool = True) -> None:
110    """Download data via https.
111
112    Args:
113        path: The path for saving the data.
114        url: The url of the data.
115        download: Whether to download the data if it is not saved at `path` yet.
116        checksum: The expected checksum of the data.
117        verify: Whether to verify the https address.
118    """
119    if os.path.exists(path):
120        return
121    if not download:
122        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
123
124    # The data is downloaded to a temporary path and only moved to `path` once it is complete and verified.
125    # Otherwise an interrupted download would be mistaken for a complete one by the check above.
126    tmp_path = f"{path}.incomplete"
127    with requests.get(url, stream=True, allow_redirects=True, verify=verify) as r:
128        r.raise_for_status()  # check for error
129        # Compute checksums on the file content rather than its HTTP transfer encoding.
130        r.raw.decode_content = True
131        file_size = int(r.headers.get("Content-Length", 0))
132        desc = f"Download {url} to {path}"
133        if file_size == 0:
134            desc += " (unknown file size)"
135        with tqdm.wrapattr(r.raw, "read", total=file_size, desc=desc) as r_raw, open(tmp_path, "wb") as f:
136            copyfileobj(r_raw, f)
137
138    _check_checksum(tmp_path, checksum)
139    os.replace(tmp_path, path)

Download data via https.

Arguments:
  • path: The path for saving the data.
  • url: The url of the data.
  • download: Whether to download the data if it is not saved at path yet.
  • checksum: The expected checksum of the data.
  • verify: Whether to verify the https address.
def download_source_gdrive( path: str, url: str, download: bool, checksum: Optional[str] = None, download_type: Literal['zip', 'folder'] = 'zip', expected_samples: int = 10000, quiet: bool = True) -> None:
142def download_source_gdrive(
143    path: str,
144    url: str,
145    download: bool,
146    checksum: Optional[str] = None,
147    download_type: Literal["zip", "folder"] = "zip",
148    expected_samples: int = 10000,
149    quiet: bool = True,
150) -> None:
151    """Download data from google drive.
152
153    Args:
154        path: The path for saving the data.
155        url: The url of the data.
156        download: Whether to download the data if it is not saved at `path` yet.
157        checksum: The expected checksum of the data.
158        download_type: The download type, either 'zip' or 'folder'.
159        expected_samples: The maximal number of samples in the folder.
160        quiet: Whether to download quietly.
161    """
162    if os.path.exists(path):
163        return
164
165    if not download:
166        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
167
168    if gdown is None:
169        raise RuntimeError(
170            "Need gdown library to download data from google drive. "
171            "Please install gdown: 'conda install -c conda-forge gdown==4.6.3'."
172        )
173
174    print("Downloading the files. Might take a few minutes...")
175
176    if download_type == "zip":
177        gdown.download(url, path, quiet=quiet)
178        _check_checksum(path, checksum)
179    elif download_type == "folder":
180        assert version.parse(gdown.__version__) == version.parse("4.6.3"), "Please install 'gdown==4.6.3'."
181        gdown.download_folder.__globals__["MAX_NUMBER_FILES"] = expected_samples
182        gdown.download_folder(url=url, output=path, quiet=quiet, remaining_ok=True)
183    else:
184        raise ValueError("`download_path` argument expects either `zip`/`folder`")
185
186    print("Download completed.")

Download data from google drive.

Arguments:
  • path: The path for saving the data.
  • url: The url of the data.
  • download: Whether to download the data if it is not saved at path yet.
  • checksum: The expected checksum of the data.
  • download_type: The download type, either 'zip' or 'folder'.
  • expected_samples: The maximal number of samples in the folder.
  • quiet: Whether to download quietly.
def download_source_empiar(path: str, access_id: str, download: bool) -> str:
189def download_source_empiar(path: str, access_id: str, download: bool) -> str:
190    """Download data from EMPIAR.
191
192    Requires the ascp command from the aspera CLI.
193
194    Args:
195        path: The path for saving the data.
196        access_id: The EMPIAR accession id of the data to download.
197        download: Whether to download the data if it is not saved at `path` yet.
198
199    Returns:
200        The path to the downloaded data.
201    """
202    download_path = os.path.join(path, access_id)
203
204    if os.path.exists(download_path):
205        return download_path
206    if not download:
207        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False")
208
209    if which("ascp") is None:
210        raise RuntimeError(
211            "Need aspera-cli to download data from empiar. You can install it via 'conda install -c hcc aspera-cli'."
212        )
213
214    key_file = os.path.expanduser("~/.aspera/cli/etc/asperaweb_id_dsa.openssh")
215    if not os.path.exists(key_file):
216        conda_root = os.environ["CONDA_PREFIX"]
217        key_file = os.path.join(conda_root, "etc/asperaweb_id_dsa.openssh")
218
219    if not os.path.exists(key_file):
220        raise RuntimeError("Could not find the aspera ssh keyfile")
221
222    cmd = ["ascp", "-QT", "-l", "200M", "-P33001", "-i", key_file, f"emp_ext2@fasp.ebi.ac.uk:/{access_id}", path]
223    run(cmd)
224
225    return download_path

Download data from EMPIAR.

Requires the ascp command from the aspera CLI.

Arguments:
  • path: The path for saving the data.
  • access_id: The EMPIAR accession id of the data to download.
  • download: Whether to download the data if it is not saved at path yet.
Returns:

The path to the downloaded data.

def download_source_kaggle( path: str, dataset_name: str, download: bool, competition: bool = False):
228def download_source_kaggle(path: str, dataset_name: str, download: bool, competition: bool = False):
229    """Download data from Kaggle.
230
231    Requires the Kaggle API.
232
233    Args:
234        path: The path for saving the data.
235        dataset_name: The name of the dataset to download.
236        download: Whether to download the data if it is not saved at `path` yet.
237        competition: Whether this data is from a competition and requires the kaggle.competition API.
238    """
239    if not download:
240        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
241
242    try:
243        from kaggle.api.kaggle_api_extended import KaggleApi
244    except ModuleNotFoundError:
245        msg = "Please install the Kaggle API. You can do this using 'pip install kaggle'. "
246        msg += "After you have installed kaggle, you would need an API token. "
247        msg += "Follow the instructions at https://www.kaggle.com/docs/api."
248        raise ModuleNotFoundError(msg)
249
250    api = KaggleApi()
251    api.authenticate()
252
253    if competition:
254        api.competition_download_files(competition=dataset_name, path=path, quiet=False)
255    else:
256        api.dataset_download_files(dataset=dataset_name, path=path, quiet=False)

Download data from Kaggle.

Requires the Kaggle API.

Arguments:
  • path: The path for saving the data.
  • dataset_name: The name of the dataset to download.
  • download: Whether to download the data if it is not saved at path yet.
  • competition: Whether this data is from a competition and requires the kaggle.competition API.
NBIA_API_URL = 'https://services.cancerimagingarchive.net/nbia-api/services/v1/'
def download_tcia_series(series_uids: List[str], dst: str, csv_filename: str) -> str:
363def download_tcia_series(series_uids: List[str], dst: str, csv_filename: str) -> str:
364    """Download individual DICOM series from TCIA by their series instance UIDs.
365
366    Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise.
367    Each series is stored in '<dst>/<SeriesInstanceUID>/', series that exist there already are skipped.
368
369    Args:
370        series_uids: The UIDs of the series to download.
371        dst: The folder for saving the DICOM series.
372        csv_filename: The path for saving the series metadata (without the '.csv' extension).
373
374    Returns:
375        The path to the csv file with the series metadata.
376    """
377    if nbia is None:
378        _download_tcia_series_with_rest(series_uids, dst, csv_filename)
379    else:
380        nbia.downloadSeries(series_data=series_uids, input_type="list", path=dst, csv_filename=csv_filename)
381    return f"{csv_filename}.csv"

Download individual DICOM series from TCIA by their series instance UIDs.

Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise. Each series is stored in '//', series that exist there already are skipped.

Arguments:
  • series_uids: The UIDs of the series to download.
  • dst: The folder for saving the DICOM series.
  • csv_filename: The path for saving the series metadata (without the '.csv' extension).
Returns:

The path to the csv file with the series metadata.

def download_source_tcia(path, url, dst, csv_filename, download):
384def download_source_tcia(path, url, dst, csv_filename, download):
385    """Download data from TCIA.
386
387    Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise.
388
389    Args:
390        path: The path for saving the manifest file. If `url` is None, this must point to an existing manifest,
391            e.g. one that was written by the caller to download only a subset of the series of a collection.
392        url: The URL to the TCIA manifest of the dataset. Set to None to use the manifest at `path`.
393        dst: The folder for saving the DICOM series. Each series is stored in a sub-folder named after its UID.
394        csv_filename: The path for saving the series metadata (without the '.csv' extension).
395        download: Whether to download the data if it is not saved at `path` yet.
396    """
397    if not download:
398        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
399
400    if url is None:
401        assert os.path.exists(path), f"The manifest {path} does not exist."
402    else:
403        assert url.endswith(".tcia"), f"{url} is not a TCIA Manifest."
404        # Downloads the manifest file from the collection page.
405        manifest = requests.get(url=url)
406        manifest.raise_for_status()
407        with open(path, "wb") as f:
408            f.write(manifest.content)
409
410    # This part extracts the UIDs from the manifests and downloads them.
411    if nbia is None:
412        _download_tcia_manifest_with_rest(path, dst, csv_filename)
413    else:
414        nbia.downloadSeries(series_data=path, input_type="manifest", path=dst, csv_filename=csv_filename)

Download data from TCIA.

Uses the tcia_utils python package if it is installed and falls back to the NBIA REST API otherwise.

Arguments:
  • path: The path for saving the manifest file. If url is None, this must point to an existing manifest, e.g. one that was written by the caller to download only a subset of the series of a collection.
  • url: The URL to the TCIA manifest of the dataset. Set to None to use the manifest at path.
  • dst: The folder for saving the DICOM series. Each series is stored in a sub-folder named after its UID.
  • csv_filename: The path for saving the series metadata (without the '.csv' extension).
  • download: Whether to download the data if it is not saved at path yet.
def download_source_synapse(path: str, entity: str, download: bool) -> None:
417def download_source_synapse(path: str, entity: str, download: bool) -> None:
418    """Download data from synapse.
419
420    Requires the synapseclient python library.
421
422    Args:
423        path: The path for saving the data.
424        entity: The name of the data to download from synapse.
425        download: Whether to download the data if it is not saved at `path` yet.
426    """
427    if not download:
428        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
429
430    if synapseclient is None:
431        raise RuntimeError(
432            "You must install 'synapseclient' to download files from 'synapse'. "
433            "Remember to create an account and generate an authentication code for your account. "
434            "Please follow the documentation for details on creating the '~/.synapseConfig' file here: "
435            "https://python-docs.synapse.org/tutorials/authentication/."
436        )
437
438    assert entity.startswith("syn"), "The entity name does not look as expected. It should be something like 'syn123'."
439
440    # Download all files in the folder.
441    syn = synapseclient.Synapse()
442    syn.login()  # Since we do not pass any credentials here, it fetches all details from '~/.synapseConfig'.
443    synapseutils.syncFromSynapse(syn=syn, entity=entity, path=path)

Download data from synapse.

Requires the synapseclient python library.

Arguments:
  • path: The path for saving the data.
  • entity: The name of the data to download from synapse.
  • download: Whether to download the data if it is not saved at path yet.
def unzip_tarfile(tar_path: str, dst: str, remove: bool = True) -> None:
456def unzip_tarfile(tar_path: str, dst: str, remove: bool = True) -> None:
457    """Unpack a tar archive.
458
459    Args:
460        tar_path: Path to the tar file.
461        dst: Where to unpack the archive.
462        remove: Whether to remove the tar file after unpacking.
463    """
464    import tarfile
465
466    if tar_path.endswith(".tar.gz") or tar_path.endswith(".tgz"):
467        access_mode = "r:gz"
468    elif tar_path.endswith(".tar"):
469        access_mode = "r:"
470    else:
471        raise ValueError(f"The provided file isn't a supported archive to unpack. Please check the file: {tar_path}.")
472
473    tar = tarfile.open(tar_path, access_mode)
474    tar.extractall(dst)
475    tar.close()
476
477    if remove:
478        os.remove(tar_path)

Unpack a tar archive.

Arguments:
  • tar_path: Path to the tar file.
  • dst: Where to unpack the archive.
  • remove: Whether to remove the tar file after unpacking.
def unzip_rarfile( rar_path: str, dst: str, remove: bool = True, use_rarfile: bool = True) -> None:
481def unzip_rarfile(rar_path: str, dst: str, remove: bool = True, use_rarfile: bool = True) -> None:
482    """Unpack a rar archive.
483
484    Args:
485        rar_path: Path to the rar file.
486        dst: Where to unpack the archive.
487        remove: Whether to remove the tar file after unpacking.
488        use_rarfile: Whether to use the rarfile library or aspose.zip.
489    """
490    def _extract_with_rarfile():
491        import rarfile
492        with rarfile.RarFile(rar_path) as archive:
493            archive.extractall(path=dst)
494
495    def _extract_with_aspose():
496        import aspose.zip as az
497        with az.rar.RarArchive(rar_path) as archive:
498            archive.extract_to_directory(dst)
499
500    def _extract_with_7z():
501        if which("7z") is None:
502            raise RuntimeError("The 'p7zip' CLI is not available.")
503        run(["7z", "x", f"-o{dst}", "-y", rar_path], check=True)
504
505    extractors = [
506        ('rarfile', _extract_with_rarfile), ('aspose.zip', _extract_with_aspose), ('7z', _extract_with_7z),
507    ] if use_rarfile else [('aspose.zip', _extract_with_aspose), ('7z', _extract_with_7z)]
508
509    errors = []
510    for name, extractor in extractors:
511        try:
512            extractor()
513            break
514        except Exception as err:
515            errors.append((name, err))
516            if len(errors) < len(extractors):
517                next_name = extractors[len(errors)][0]
518                warn(f"Extraction with '{name}' failed for {rar_path} ({err}). Falling back to '{next_name}'.")
519    else:
520        backends = ', '.join(f"'{name}'" for name, _ in extractors)
521        raise RuntimeError(
522            f"Failed to extract rar archive {rar_path} with {backends}. "
523            "Please ensure one of the supported backends is installed and can read this archive."
524        ) from errors[-1][1]
525
526    if remove:
527        os.remove(rar_path)

Unpack a rar archive.

Arguments:
  • rar_path: Path to the rar file.
  • dst: Where to unpack the archive.
  • remove: Whether to remove the tar file after unpacking.
  • use_rarfile: Whether to use the rarfile library or aspose.zip.
def unzip(zip_path: str, dst: str, remove: bool = True) -> None:
530def unzip(zip_path: str, dst: str, remove: bool = True) -> None:
531    """Unpack a zip archive.
532
533    Args:
534        zip_path: Path to the zip file.
535        dst: Where to unpack the archive.
536        remove: Whether to remove the tar file after unpacking.
537    """
538    with zipfile.ZipFile(zip_path, "r") as f:
539        f.extractall(dst)
540    if remove:
541        os.remove(zip_path)

Unpack a zip archive.

Arguments:
  • zip_path: Path to the zip file.
  • dst: Where to unpack the archive.
  • remove: Whether to remove the tar file after unpacking.
def unzip_7z(path_7z: str, dst: str, remove: bool = True) -> None:
544def unzip_7z(path_7z: str, dst: str, remove: bool = True) -> None:
545    """Unpack a 7z archive.
546
547    Args:
548        path_7z: Path to the 7z file.
549        dst: Where to unpack the archive.
550        remove: Whether to remove the 7z file after unpacking.
551    """
552    if which("7z") is None:
553        raise RuntimeError("Need the 'p7zip' CLI to extract 7z archives. You can install it via 'conda install -c conda-forge p7zip'.")  # noqa
554
555    run(["7z", "x", f"-o{dst}", "-y", path_7z])
556
557    if remove:
558        os.remove(path_7z)

Unpack a 7z archive.

Arguments:
  • path_7z: Path to the 7z file.
  • dst: Where to unpack the archive.
  • remove: Whether to remove the 7z file after unpacking.
def generate_labeled_array_from_xml(shape: Tuple[int, ...], xml_file: str) -> numpy.ndarray:
657def generate_labeled_array_from_xml(shape: Tuple[int, ...], xml_file: str) -> np.ndarray:
658    """Generate a label mask from a contour defined in a xml annotation file.
659
660    Function taken from: https://github.com/rshwndsz/hover-net/blob/master/lightning_hovernet.ipynb
661
662    Args:
663        shape: The image shape.
664        xml_file: The path to the xml file with contour annotations.
665
666    Returns:
667        The label mask.
668    """
669    # DOM object created by the minidom parser
670    xDoc = minidom.parse(xml_file)
671
672    # List of all Region tags
673    regions = xDoc.getElementsByTagName('Region')
674
675    # List which will store the vertices for each region
676    xy = []
677    for region in regions:
678        # Loading all the vertices in the region
679        vertices = region.getElementsByTagName('Vertex')
680
681        # The vertices of a region will be stored in a array
682        vw = np.zeros((len(vertices), 2))
683
684        for index, vertex in enumerate(vertices):
685            # Storing the values of x and y coordinate after conversion
686            vw[index][0] = float(vertex.getAttribute('X'))
687            vw[index][1] = float(vertex.getAttribute('Y'))
688
689        # Append the vertices of a region
690        xy.append(np.int32(vw))
691
692    # Creating a completely black image
693    mask = np.zeros(shape, np.uint32)  # Integer instance ids; float labels break connected-component ops.
694
695    # Start the instance ids at 1: id 0 is background, so enumerating from 0 silently drops the first region.
696    for i, contour in enumerate(xy, start=1):
697        r, c = polygon(np.array(contour)[:, 1], np.array(contour)[:, 0], shape=shape)
698        mask[r, c] = i
699    return mask

Generate a label mask from a contour defined in a xml annotation file.

Function taken from: https://github.com/rshwndsz/hover-net/blob/master/lightning_hovernet.ipynb

Arguments:
  • shape: The image shape.
  • xml_file: The path to the xml file with contour annotations.
Returns:

The label mask.

def load_dicom_series(series_dir: str) -> Tuple[numpy.ndarray, Dict[str, numpy.ndarray]]:
702def load_dicom_series(series_dir: str) -> Tuple[np.ndarray, Dict[str, np.ndarray]]:
703    """Stack a single-frame DICOM image series (CT, MR, PET) into a volume with axes (z, y, x).
704
705    The slices are sorted by their position along the slice normal (the cross product of the row and column
706    direction in 'ImageOrientationPatient'), so the volume is stacked consistently for any acquisition plane.
707    'RescaleSlope' and 'RescaleIntercept' are applied per slice, i.e. CT volumes are returned in Hounsfield units.
708
709    NOTE: This requires the pydicom python package.
710
711    Args:
712        series_dir: The folder with the DICOM files of the series.
713
714    Returns:
715        The volume with axes (z, y, x) as float32.
716        The geometry of the volume, which is needed by `rasterize_rtstruct`. A dictionary with the keys
717        'origin' (the 'ImagePositionPatient' of each slice, n_slices x 3), 'row_direction' and 'column_direction'
718        (the unit vectors along which the column and the row index increase, from 'ImageOrientationPatient'),
719        'spacing' (the row and column spacing from 'PixelSpacing') and 'sop_uids' (the 'SOPInstanceUID' per slice).
720    """
721    import pydicom
722
723    dcm_paths = [os.path.join(series_dir, fname) for fname in sorted(os.listdir(series_dir)) if fname.endswith(".dcm")]
724    slices = [pydicom.dcmread(dcm_path) for dcm_path in dcm_paths]
725
726    row_direction = np.array([float(v) for v in slices[0].ImageOrientationPatient[:3]])
727    column_direction = np.array([float(v) for v in slices[0].ImageOrientationPatient[3:]])
728    normal = np.cross(row_direction, column_direction)
729    slices.sort(key=lambda dcm: np.dot([float(v) for v in dcm.ImagePositionPatient], normal))
730
731    volume = []
732    for dcm in slices:
733        frame = dcm.pixel_array.astype("float32")
734        volume.append(frame * float(dcm.get("RescaleSlope", 1.0)) + float(dcm.get("RescaleIntercept", 0.0)))
735    volume = np.stack(volume)
736
737    geometry = {
738        "origin": np.array([[float(v) for v in dcm.ImagePositionPatient] for dcm in slices]),
739        "row_direction": row_direction,
740        "column_direction": column_direction,
741        "spacing": np.array([float(v) for v in slices[0].PixelSpacing]),
742        "sop_uids": np.array([str(dcm.SOPInstanceUID) for dcm in slices]),
743    }
744    return volume, geometry

Stack a single-frame DICOM image series (CT, MR, PET) into a volume with axes (z, y, x).

The slices are sorted by their position along the slice normal (the cross product of the row and column direction in 'ImageOrientationPatient'), so the volume is stacked consistently for any acquisition plane. 'RescaleSlope' and 'RescaleIntercept' are applied per slice, i.e. CT volumes are returned in Hounsfield units.

NOTE: This requires the pydicom python package.

Arguments:
  • series_dir: The folder with the DICOM files of the series.
Returns:

The volume with axes (z, y, x) as float32. The geometry of the volume, which is needed by rasterize_rtstruct. A dictionary with the keys 'origin' (the 'ImagePositionPatient' of each slice, n_slices x 3), 'row_direction' and 'column_direction' (the unit vectors along which the column and the row index increase, from 'ImageOrientationPatient'), 'spacing' (the row and column spacing from 'PixelSpacing') and 'sop_uids' (the 'SOPInstanceUID' per slice).

def rasterize_rtstruct( rtstruct_path: str, geometry: Dict[str, numpy.ndarray], shape: Tuple[int, int, int], roi_labels: Union[Dict[str, int], Callable[[int, str], Optional[int]]]) -> numpy.ndarray:
747def rasterize_rtstruct(
748    rtstruct_path: str,
749    geometry: Dict[str, np.ndarray],
750    shape: Tuple[int, int, int],
751    roi_labels: Union[Dict[str, int], Callable[[int, str], Optional[int]]],
752) -> np.ndarray:
753    """Rasterize the contours of a DICOM RTSTRUCT file onto the voxel grid of the referenced image series.
754
755    Each 'CLOSED_PLANAR' contour is assigned to the slice it references (via 'ReferencedSOPInstanceUID', with a
756    fallback to the closest slice along the slice normal). Its points are projected onto the row and column
757    direction of that slice to obtain pixel coordinates and the polygon is filled with `skimage.draw.polygon`.
758    Multiple contours of the same ROI on the same slice are combined with XOR, so that inner contours form holes.
759    Where different ROIs overlap, the ROI with the lower label id takes precedence.
760
761    NOTE: This requires the pydicom python package.
762
763    Args:
764        rtstruct_path: The path to the RTSTRUCT DICOM file.
765        geometry: The geometry of the referenced image series, as returned by `load_dicom_series`.
766        shape: The shape of the image volume (z, y, x).
767        roi_labels: The mapping from ROIs to label ids. Either a dictionary that maps the ROI names to label ids,
768            or a function that maps the ROI number and ROI name to a label id. ROIs that are not in the dictionary
769            or for which the function returns None are ignored.
770
771    Returns:
772        The label volume (uint8) with axes (z, y, x).
773    """
774    import pydicom
775
776    rtstruct = pydicom.dcmread(rtstruct_path)
777    roi_names = {int(roi.ROINumber): str(roi.ROIName) for roi in rtstruct.StructureSetROISequence}
778    slice_ids = {uid: z for z, uid in enumerate(geometry["sop_uids"])}
779    normal = np.cross(geometry["row_direction"], geometry["column_direction"])
780    slice_positions = geometry["origin"] @ normal
781
782    masks = {}
783    for roi_contour in rtstruct.ROIContourSequence:
784        roi_number = int(roi_contour.ReferencedROINumber)
785        roi_name = roi_names[roi_number]
786        if isinstance(roi_labels, dict):
787            label_id = roi_labels.get(roi_name)
788        else:
789            label_id = roi_labels(roi_number, roi_name)
790        if label_id is None:
791            continue
792
793        mask = masks.setdefault(label_id, np.zeros(shape, dtype="bool"))
794        for contour in roi_contour.get("ContourSequence", []):
795            if contour.ContourGeometricType != "CLOSED_PLANAR":
796                continue
797            points = np.array([float(v) for v in contour.ContourData]).reshape(-1, 3)
798            if len(points) < 3:
799                continue
800
801            z = None
802            for ref in contour.get("ContourImageSequence", []):
803                z = slice_ids.get(str(ref.ReferencedSOPInstanceUID))
804            if z is None:
805                z = int(np.argmin(np.abs(slice_positions - np.dot(points[0], normal))))
806
807            offsets = points - geometry["origin"][z]
808            cols = offsets @ geometry["row_direction"] / geometry["spacing"][1]
809            rows = offsets @ geometry["column_direction"] / geometry["spacing"][0]
810            rr, cc = polygon(rows, cols, shape=shape[1:])
811            mask[z, rr, cc] = ~mask[z, rr, cc]
812
813    labels = np.zeros(shape, dtype="uint8")
814    for label_id in sorted(masks, reverse=True):
815        labels[masks[label_id]] = label_id
816    return labels

Rasterize the contours of a DICOM RTSTRUCT file onto the voxel grid of the referenced image series.

Each 'CLOSED_PLANAR' contour is assigned to the slice it references (via 'ReferencedSOPInstanceUID', with a fallback to the closest slice along the slice normal). Its points are projected onto the row and column direction of that slice to obtain pixel coordinates and the polygon is filled with skimage.draw.polygon. Multiple contours of the same ROI on the same slice are combined with XOR, so that inner contours form holes. Where different ROIs overlap, the ROI with the lower label id takes precedence.

NOTE: This requires the pydicom python package.

Arguments:
  • rtstruct_path: The path to the RTSTRUCT DICOM file.
  • geometry: The geometry of the referenced image series, as returned by load_dicom_series.
  • shape: The shape of the image volume (z, y, x).
  • roi_labels: The mapping from ROIs to label ids. Either a dictionary that maps the ROI names to label ids, or a function that maps the ROI number and ROI name to a label id. ROIs that are not in the dictionary or for which the function returns None are ignored.
Returns:

The label volume (uint8) with axes (z, y, x).

def convert_svs_to_array( path: str, location: Tuple[int, int] = (0, 0), level: int = 0, img_size: Tuple[int, int] = None) -> numpy.ndarray:
820def convert_svs_to_array(
821    path: str, location: Tuple[int, int] = (0, 0), level: int = 0, img_size: Tuple[int, int] = None,
822) -> np.ndarray:
823    """Convert a .svs file for WSI imagging to a numpy array.
824
825    Requires the tiffslide python library.
826    The function can load multi-resolution images. You can specify the resolution level via `level`.
827
828    Args:
829        path: File path ath to the svs file.
830        location: Pixel location (x, y) in level 0 of the image.
831        level: Target level used to read the image.
832        img_size: Size of the image. If None, the shape of the image at `level` is used.
833
834    Returns:
835        The image as numpy array.
836    """
837    assert path.endswith(".svs"), f"The provided file ({path}) isn't in svs format"
838
839    try:
840        from tiffslide import TiffSlide
841    except ImportError:
842        # svs is a pyramidal TIFF variant, so tifffile can read it without the tiffslide dependency.
843        import tifffile
844        with tifffile.TiffFile(path) as f:
845            image = f.series[0].levels[level].asarray()
846        x, y = location
847        if img_size is not None:
848            image = image[y:y + img_size[1], x:x + img_size[0]]
849        else:
850            image = image[y:, x:]
851        return image
852
853    _slide = TiffSlide(path)
854    if img_size is None:
855        img_size = _slide.level_dimensions[0]
856    return _slide.read_region(location=location, level=level, size=img_size, as_array=True)

Convert a .svs file for WSI imagging to a numpy array.

Requires the tiffslide python library. The function can load multi-resolution images. You can specify the resolution level via level.

Arguments:
  • path: File path ath to the svs file.
  • location: Pixel location (x, y) in level 0 of the image.
  • level: Target level used to read the image.
  • img_size: Size of the image. If None, the shape of the image at level is used.
Returns:

The image as numpy array.

def download_from_cryo_et_portal(path: str, dataset_id: int, download: bool) -> str:
859def download_from_cryo_et_portal(path: str, dataset_id: int, download: bool) -> str:
860    """Download data from the CryoET Data Portal.
861
862    Requires the cryoet-data-portal python library.
863
864    Args:
865        path: The path for saving the data.
866        dataset_id: The id of the data to download from the portal.
867        download: Whether to download the data if it is not saved at `path` yet.
868
869    Returns:
870        The file path to the downloaded data.
871    """
872    if Client is None or Dataset is None:
873        raise RuntimeError("Please install CryoETDataPortal via 'pip install cryoet-data-portal'")
874
875    output_path = os.path.join(path, str(dataset_id))
876    if os.path.exists(output_path):
877        return output_path
878
879    if not download:
880        raise RuntimeError(f"Cannot find the data at {path}, but download was set to False.")
881
882    client = Client()
883    dataset = Dataset.get_by_id(client, dataset_id)
884    dataset.download_everything(dest_path=path)
885
886    return output_path

Download data from the CryoET Data Portal.

Requires the cryoet-data-portal python library.

Arguments:
  • path: The path for saving the data.
  • dataset_id: The id of the data to download from the portal.
  • download: Whether to download the data if it is not saved at path yet.
Returns:

The file path to the downloaded data.