feat: GPU detection + PHOTO_JUDGERS_DEVICE override + diagnostics
This commit is contained in:
+20
-1
@@ -68,7 +68,26 @@ class LaionScorer:
|
||||
self.head = None # type: ignore[assignment]
|
||||
self._version = "v2"
|
||||
self._cache_dir = cache_dir
|
||||
self._device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
self._device = self._pick_device()
|
||||
|
||||
@staticmethod
|
||||
def _pick_device() -> str:
|
||||
"""Select CUDA GPU when available; allow override via env var.
|
||||
|
||||
Priority: PHOTO_JUDGERS_DEVICE env var -> cuda (if available) -> cpu.
|
||||
"""
|
||||
override = os.environ.get("PHOTO_JUDGERS_DEVICE", "").strip().lower()
|
||||
if override:
|
||||
print(f" Device overridden by PHOTO_JUDGERS_DEVICE: {override}")
|
||||
return override
|
||||
|
||||
cuda_ok = torch.cuda.is_available()
|
||||
print(f" torch.cuda.is_available(): {cuda_ok}")
|
||||
if cuda_ok:
|
||||
print(f" Using GPU: {torch.cuda.get_device_name(0)}")
|
||||
return "cuda"
|
||||
print(" Using CPU (CUDA not available to torch).")
|
||||
return "cpu"
|
||||
|
||||
def _head_path(self) -> str:
|
||||
return os.path.join(self._cache_dir, self.HEAD_FILENAME)
|
||||
|
||||
Reference in New Issue
Block a user