File size: 1,052 Bytes
a358495
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
"""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"]