mirror of
https://github.com/snapotter-hq/SnapOtter.git
synced 2026-08-03 07:46:42 +02:00
fix(ai): use centralized GPU detection in enhance_faces and inpaint (#63)
enhance_faces.py relied on implicit PyTorch auto-detection for both GFPGAN and CodeFormer, bypassing the centralized gpu.py module. inpaint.py queried ort.get_available_providers() directly, which reports compiled-in backends rather than actual hardware. Both tools now go through gpu.py so STIRLING_GPU=false correctly forces CPU across every AI tool. Co-authored-by: stirling-image <stirling-image@users.noreply.github.com>
This commit is contained in:
co-authored by
stirling-image
parent
6a43cc1b77
commit
43821a955c
@@ -80,17 +80,23 @@ def detect_faces_mediapipe(img_array, sensitivity):
|
|||||||
|
|
||||||
def enhance_with_gfpgan(img_array, only_center_face):
|
def enhance_with_gfpgan(img_array, only_center_face):
|
||||||
"""Enhance faces using GFPGAN. Returns the enhanced image array."""
|
"""Enhance faces using GFPGAN. Returns the enhanced image array."""
|
||||||
|
import torch
|
||||||
from gfpgan import GFPGANer
|
from gfpgan import GFPGANer
|
||||||
|
from gpu import gpu_available
|
||||||
|
|
||||||
if not os.path.exists(GFPGAN_MODEL_PATH):
|
if not os.path.exists(GFPGAN_MODEL_PATH):
|
||||||
raise FileNotFoundError(f"GFPGAN model not found: {GFPGAN_MODEL_PATH}")
|
raise FileNotFoundError(f"GFPGAN model not found: {GFPGAN_MODEL_PATH}")
|
||||||
|
|
||||||
|
use_gpu = gpu_available()
|
||||||
|
device = torch.device("cuda" if use_gpu else "cpu")
|
||||||
|
|
||||||
enhancer = GFPGANer(
|
enhancer = GFPGANer(
|
||||||
model_path=GFPGAN_MODEL_PATH,
|
model_path=GFPGAN_MODEL_PATH,
|
||||||
upscale=1,
|
upscale=1,
|
||||||
arch="clean",
|
arch="clean",
|
||||||
channel_multiplier=2,
|
channel_multiplier=2,
|
||||||
bg_upsampler=None,
|
bg_upsampler=None,
|
||||||
|
device=device,
|
||||||
)
|
)
|
||||||
_, _, output = enhancer.enhance(
|
_, _, output = enhancer.enhance(
|
||||||
img_array,
|
img_array,
|
||||||
@@ -115,23 +121,33 @@ def enhance_with_codeformer(img_array, fidelity_weight):
|
|||||||
fails, the auto model selection will fall back to GFPGAN.
|
fails, the auto model selection will fall back to GFPGAN.
|
||||||
"""
|
"""
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from gpu import gpu_available
|
||||||
|
|
||||||
# Import may fail if codeformer-pip is not installed or if the
|
use_gpu = gpu_available()
|
||||||
# module-level model loading fails (missing weights, no GPU, etc.)
|
|
||||||
from codeformer.app import inference_app
|
# CodeFormer selects its device during module-level init and inside
|
||||||
|
# inference_app(). It has no device= parameter, so to respect
|
||||||
|
# STIRLING_GPU=false we temporarily override torch.cuda.is_available
|
||||||
|
# so all internal device checks see False. When use_gpu is True
|
||||||
|
# (the common path) no override happens.
|
||||||
|
_orig_cuda_check = torch.cuda.is_available
|
||||||
|
if not use_gpu:
|
||||||
|
torch.cuda.is_available = lambda: False
|
||||||
|
try:
|
||||||
|
from codeformer.app import inference_app
|
||||||
|
|
||||||
|
img_bgr = img_array[:, :, ::-1].copy()
|
||||||
|
restored_bgr = inference_app(
|
||||||
|
image=img_bgr,
|
||||||
|
background_enhance=False,
|
||||||
|
face_upsample=False,
|
||||||
|
upscale=1,
|
||||||
|
codeformer_fidelity=fidelity_weight,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
torch.cuda.is_available = _orig_cuda_check
|
||||||
|
|
||||||
# inference_app accepts a numpy array (BGR) or file path.
|
|
||||||
# It returns the restored image as a BGR numpy array.
|
|
||||||
# We pass our RGB array converted to BGR since OpenCV convention is used internally.
|
|
||||||
img_bgr = img_array[:, :, ::-1].copy()
|
|
||||||
restored_bgr = inference_app(
|
|
||||||
image=img_bgr,
|
|
||||||
background_enhance=False,
|
|
||||||
face_upsample=False,
|
|
||||||
upscale=1,
|
|
||||||
codeformer_fidelity=fidelity_weight,
|
|
||||||
)
|
|
||||||
# Convert back to RGB
|
|
||||||
restored_rgb = restored_bgr[:, :, ::-1].copy()
|
restored_rgb = restored_bgr[:, :, ::-1].copy()
|
||||||
return restored_rgb
|
return restored_rgb
|
||||||
|
|
||||||
|
|||||||
@@ -110,9 +110,8 @@ def main():
|
|||||||
model_path = _get_model_path()
|
model_path = _get_model_path()
|
||||||
|
|
||||||
# Configure ONNX Runtime session
|
# Configure ONNX Runtime session
|
||||||
providers = ["CPUExecutionProvider"]
|
from gpu import onnx_providers
|
||||||
if "CUDAExecutionProvider" in ort.get_available_providers():
|
providers = onnx_providers()
|
||||||
providers.insert(0, "CUDAExecutionProvider")
|
|
||||||
|
|
||||||
session = ort.InferenceSession(model_path, providers=providers)
|
session = ort.InferenceSession(model_path, providers=providers)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user