"""
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