SM737's picture
Upload folder using huggingface_hub (part 4)
a358495 verified
Raw
History Blame Contribute Delete
1.05 kB
"""One runtime device policy for the active PyTorch and ONNX model routes."""
from __future__ import annotations
import os
def torch_device() -> str:
import torch
requested = os.getenv("SATQUERY_DEVICE", "auto").lower()
if requested not in {"auto", "cuda", "cpu"}:
raise ValueError("SATQUERY_DEVICE must be auto, cuda, or cpu")
if requested == "cpu":
return "cpu"
if torch.cuda.is_available():
return "cuda"
if requested == "cuda":
raise RuntimeError("CUDA was requested but PyTorch cannot access the GPU")
return "cpu"
def onnx_providers() -> list[str]:
import onnxruntime as ort
if torch_device() == "cuda":
if "CUDAExecutionProvider" not in ort.get_available_providers():
raise RuntimeError("CUDA is available, but ONNX Runtime has no CUDA execution provider")
try:
ort.preload_dlls()
except AttributeError:
pass
return ["CUDAExecutionProvider", "CPUExecutionProvider"]
return ["CPUExecutionProvider"]