Source code for medical_image.utils.device

"""
GPU device management, memory handling, mixed precision, and multi-GPU utilities.
"""

import functools
import threading
from enum import Enum
from typing import List, Optional, Union

import torch

from medical_image.utils.logging import logger

# ---------------------------------------------------------------------------
# Device resolution
# ---------------------------------------------------------------------------


[docs] def resolve_device( *images, explicit: Union[str, torch.device, None] = None ) -> torch.device: """ Determine the target device for a processing operation. Priority: 1. Explicit device parameter (if provided) 2. Device of the first loaded image 3. Fallback to CPU """ if explicit is not None: return torch.device(explicit) for img in images: if hasattr(img, "pixel_data") and img.pixel_data is not None: return img.pixel_data.device return torch.device("cpu")
# --------------------------------------------------------------------------- # Mixed precision # ---------------------------------------------------------------------------
[docs] class Precision(Enum): FULL = torch.float32 HALF = torch.float16 BFLOAT16 = torch.bfloat16
_default_precision: Precision = Precision.FULL _precision_lock = threading.Lock()
[docs] def set_default_precision(precision: Precision) -> None: global _default_precision with _precision_lock: _default_precision = precision
[docs] def get_default_precision() -> Precision: with _precision_lock: return _default_precision
[docs] def get_dtype() -> torch.dtype: with _precision_lock: return _default_precision.value
# --------------------------------------------------------------------------- # DeviceContext — GPU-aware context manager # ---------------------------------------------------------------------------
[docs] class DeviceContext: """ Context manager for GPU-aware processing with automatic memory management. Features: - Clears GPU cache on entry and exit - Provides memory usage tracking - Automatic CPU fallback when CUDA is unavailable """
[docs] def __init__( self, device: str = "cuda", fallback: str = "cpu", verbose: bool = False, ): self.primary = torch.device(device) self.fallback = torch.device(fallback) self.active_device = self.primary self.verbose = verbose
def __enter__(self) -> "DeviceContext": if self.primary.type == "cuda" and torch.cuda.is_available(): torch.cuda.empty_cache() if self.verbose: free, total = torch.cuda.mem_get_info(self.primary) logger.info(f"GPU memory: {free / 1e9:.1f} / {total / 1e9:.1f} GB free") elif self.primary.type == "cuda": logger.warning("CUDA requested but not available — falling back to CPU") self.active_device = self.fallback return self def __exit__(self, exc_type, exc_val, exc_tb): if self.active_device.type == "cuda": torch.cuda.empty_cache() if exc_type is torch.cuda.OutOfMemoryError: logger.error(f"GPU OOM in DeviceContext: {exc_val}") self.active_device = self.fallback # Don't suppress — caller must handle retry explicitly return False return False @property def device(self) -> torch.device: return self.active_device
[docs] def memory_stats(self) -> dict: """Return current GPU memory usage.""" if self.active_device.type != "cuda": return {"device": "cpu"} free, total = torch.cuda.mem_get_info(self.active_device) return { "device": str(self.active_device), "allocated_gb": torch.cuda.memory_allocated(self.active_device) / 1e9, "free_gb": free / 1e9, "total_gb": total / 1e9, }
# --------------------------------------------------------------------------- # @gpu_safe — OOM fallback decorator # ---------------------------------------------------------------------------
[docs] def gpu_safe(func): """Decorator: catches CUDA errors (OOM, device-side asserts) and retries on CPU.""" @functools.wraps(func) def wrapper(*args, device=None, **kwargs): device = resolve_device( *[a for a in args if hasattr(a, "pixel_data")], explicit=device ) try: return func(*args, device=device, **kwargs) except (torch.cuda.OutOfMemoryError, RuntimeError, torch.AcceleratorError) as e: if device.type != "cpu" and ( isinstance(e, (torch.cuda.OutOfMemoryError, torch.AcceleratorError)) or "CUDA" in str(e) ): logger.warning(f"{func.__name__}: CUDA error — retrying on CPU: {e}") torch.cuda.empty_cache() return func(*args, device=torch.device("cpu"), **kwargs) raise return wrapper
# --------------------------------------------------------------------------- # AsyncGPUPipeline — overlapped I/O + compute with CUDA streams # ---------------------------------------------------------------------------
[docs] class AsyncGPUPipeline: """ Overlap disk I/O, CPU→GPU transfer, and GPU compute using CUDA streams. Only usable when CUDA is available. """
[docs] def __init__(self, device: str = "cuda"): self.device = torch.device(device) if not torch.cuda.is_available(): raise RuntimeError("AsyncGPUPipeline requires CUDA") self.compute_stream = torch.cuda.Stream(self.device) self.transfer_stream = torch.cuda.Stream(self.device)
[docs] def process_images(self, images: list, algorithm) -> list: """ Process pre-loaded Image objects with overlapped transfer and compute. Args: images: List of Image objects (already loaded). algorithm: An Algorithm instance. Returns: List of output Image objects. """ results = [] for img in images: # Transfer to GPU on transfer stream with torch.cuda.stream(self.transfer_stream): gpu_data = img.pixel_data.pin_memory().to( self.device, non_blocking=True ) # Compute on compute stream — don't mutate original image with torch.cuda.stream(self.compute_stream): self.compute_stream.wait_stream(self.transfer_stream) gpu_img = img.clone() gpu_img.pixel_data = gpu_data output = gpu_img.clone() algorithm(gpu_img, output) results.append(output) torch.cuda.synchronize() return results
# --------------------------------------------------------------------------- # MultiGPUAlgorithm — data-parallel across GPUs # ---------------------------------------------------------------------------
[docs] def check_gpu_budget(required_bytes: int, device: torch.device = None) -> bool: """Return True if enough GPU memory is available for the operation. Args: required_bytes: Estimated memory needed in bytes. device: Target CUDA device. Returns True for non-CUDA devices. """ if device is None or device.type != "cuda": return True free, _ = torch.cuda.mem_get_info(device) return free >= required_bytes
[docs] def estimate_image_bytes(image, dtype: torch.dtype = torch.float32) -> int: """Estimate GPU memory needed for an image tensor. Args: image: An Image object or any object with pixel_data/width/height. dtype: Assumed dtype if pixel_data is not loaded. Returns: Estimated bytes required. """ if hasattr(image, "pixel_data") and image.pixel_data is not None: return image.pixel_data.nelement() * image.pixel_data.element_size() if hasattr(image, "width") and hasattr(image, "height"): w, h = image.width, image.height if w and h: return w * h * torch.tensor([], dtype=dtype).element_size() return 0
# --------------------------------------------------------------------------- # MultiGPUAlgorithm — data-parallel across GPUs # ---------------------------------------------------------------------------
[docs] class MultiGPUAlgorithm: """ Distribute algorithm execution across available GPUs (data-parallel). """
[docs] def __init__( self, algorithm_cls: type, gpu_ids: Optional[List[int]] = None, **kwargs, ): if not torch.cuda.is_available(): raise RuntimeError("MultiGPUAlgorithm requires CUDA") if gpu_ids is None: gpu_ids = list(range(torch.cuda.device_count())) self.gpu_ids = gpu_ids self.algorithms = { gpu_id: algorithm_cls(device=f"cuda:{gpu_id}", **kwargs) for gpu_id in gpu_ids }
[docs] def apply_batch(self, images: list, outputs: list) -> list: """Distribute images across GPUs round-robin.""" n_gpus = len(self.gpu_ids) results = [None] * len(images) for i, (img, out) in enumerate(zip(images, outputs)): gpu_id = self.gpu_ids[i % n_gpus] algo = self.algorithms[gpu_id] img.to(f"cuda:{gpu_id}") out.to(f"cuda:{gpu_id}") algo.apply(img, out) results[i] = out return results