ai-gateway / core /runtime.py
basyx's picture
Upload 61 files
d63e724 verified
Raw
History Blame Contribute Delete
4.93 kB
"""Accelerator selection, inference contexts, and memory cleanup."""
from __future__ import annotations
import gc
import logging
from contextlib import ExitStack, contextmanager
from typing import Any, Iterator
from config import Settings
logger = logging.getLogger(__name__)
def import_torch() -> Any:
"""Import torch lazily so startup never loads a CUDA context."""
try:
import torch
except ImportError as exc: # pragma: no cover - required in deployment
raise RuntimeError("PyTorch is not installed") from exc
return torch
def detect_device(configured: str = "auto") -> str:
"""Select the requested accelerator, falling back safely to CPU."""
try:
torch = import_torch()
except RuntimeError:
if configured in {"auto", "cpu"}:
return "cpu"
raise
if configured != "auto":
if configured.startswith("cuda"):
if not torch.cuda.is_available():
return "cpu"
if ":" in configured:
device_index = int(configured.split(":", 1)[1])
if device_index >= torch.cuda.device_count():
return "cpu"
if configured == "mps":
mps = getattr(torch.backends, "mps", None)
return "mps" if mps is not None and mps.is_available() else "cpu"
return configured
if torch.cuda.is_available():
return "cuda"
mps = getattr(torch.backends, "mps", None)
if mps is not None and mps.is_available():
return "mps"
return "cpu"
def preferred_dtype(device: str, mixed_precision: bool = True) -> Any:
"""Choose an inference dtype supported by the selected device."""
torch = import_torch()
if not mixed_precision or device == "cpu":
return torch.float32
if device.startswith("cuda"):
return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
return torch.float16 if device == "mps" else torch.float32
def gpu_available() -> bool:
"""Return scheduled ZeroGPU or directly attached CUDA availability."""
if zerogpu_enabled():
return True
try:
torch = import_torch()
except RuntimeError:
return False
return bool(torch.cuda.is_available())
def zerogpu_enabled() -> bool:
"""Return whether this process is running in a Hugging Face ZeroGPU Space."""
try:
from spaces.config import Config
except ImportError:
return False
return bool(Config.zero_gpu)
def initialize_zerogpu() -> None:
"""Run the scheduler startup hook when Gradio is mounted into FastAPI."""
if not zerogpu_enabled():
return
from spaces.zero import startup
startup()
@contextmanager
def inference_context(settings: Settings, device: str) -> Iterator[None]:
"""Disable autograd and enable CUDA mixed precision when configured."""
torch = import_torch()
with ExitStack() as stack:
stack.enter_context(torch.no_grad())
stack.enter_context(torch.inference_mode())
if settings.mixed_precision and device.startswith("cuda"):
stack.enter_context(
torch.autocast(device_type="cuda", dtype=preferred_dtype(device, True))
)
yield
def cleanup_memory() -> None:
"""Collect Python garbage and return cached accelerator memory."""
gc.collect()
try:
torch = import_torch()
except RuntimeError:
return
# Cache cleanup must never initialize CUDA. ZeroGPU deliberately reports CUDA
# as available in the parent process while rejecting low-level initialization;
# a real CUDA context exists only inside the scheduled worker.
try:
if torch.cuda.is_initialized():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
except Exception:
logger.debug("CUDA cache cleanup skipped", exc_info=True)
mps_backend = getattr(getattr(torch, "backends", None), "mps", None)
mps = getattr(torch, "mps", None)
try:
if (
mps_backend is not None
and mps_backend.is_available()
and mps is not None
and hasattr(mps, "empty_cache")
):
mps.empty_cache()
except Exception:
logger.debug("MPS cache cleanup skipped", exc_info=True)
def memory_stats() -> dict[str, float]:
"""Return process and CUDA memory consumption in MiB."""
stats: dict[str, float] = {}
try:
import psutil
stats["rss_mb"] = round(psutil.Process().memory_info().rss / 1024**2, 2)
except ImportError:
pass
try:
torch = import_torch()
except RuntimeError:
return stats
if torch.cuda.is_initialized():
stats.update(
cuda_allocated_mb=round(torch.cuda.memory_allocated() / 1024**2, 2),
cuda_reserved_mb=round(torch.cuda.memory_reserved() / 1024**2, 2),
)
return stats