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
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.
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
pathyet. - checksum: The expected checksum of the data.
- verify: Whether to verify the https address.
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
pathyet. - 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.
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
pathyet.
Returns:
The path to the downloaded data.
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
pathyet. - competition: Whether this data is from a competition and requires the kaggle.competition API.
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 '
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.
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
urlis 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
pathyet.
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
pathyet.
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.
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.
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.
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.
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.
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).
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).
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
levelis used.
Returns:
The image as numpy array.
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
pathyet.
Returns:
The file path to the downloaded data.