File size: 2,753 Bytes
517b1e7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
# -*- coding: utf-8 -*-
"""Face restoration.

This is the GFPGANer pipeline — detect, align, restore, paste back — rebuilt on
facexlib's FaceRestoreHelper plus a spandrel-loaded GFPGAN checkpoint, so the
abandoned gfpgan and basicsr packages are no longer needed.
"""

from __future__ import annotations

import threading

import numpy as np

from upscale import WEIGHTS_DIR, load_face_model, run_model

_helper = None
_helper_lock = threading.Lock()


def get_helper(upscale: int):
    """Returns a process-wide FaceRestoreHelper, retuned to the given upscale.

    The helper carries per-image state (landmarks, affines, cropped faces), so
    callers must hold `helper_lock()` for the whole detect->paste sequence.
    """
    global _helper
    import torch
    from facexlib.utils.face_restoration_helper import FaceRestoreHelper

    # facexlib's detection and parsing nets are kept on CUDA or CPU only. MPS is
    # deliberately excluded — these nets are untested there and the detector is
    # cheap relative to the SR pass, which still runs on the accelerator.
    det_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    if _helper is None:
        _helper = FaceRestoreHelper(
            upscale_factor=upscale,
            face_size=512,
            crop_ratio=(1, 1),
            det_model="retinaface_resnet50",
            save_ext="png",
            use_parse=True,
            device=det_device,
            model_rootpath=WEIGHTS_DIR,
        )
    else:
        _helper.set_upscale_factor(upscale)
    return _helper


def helper_lock() -> threading.Lock:
    return _helper_lock


def restore_faces(bgr: np.ndarray, background: np.ndarray, upscale: int) -> np.ndarray:
    """Restores every detected face in `bgr` and pastes them onto `background`.

    `bgr` is the original BGR image, `background` the already-upscaled BGR image
    the faces are composited onto. Returns BGR. If no face is found, the
    background is returned untouched.
    """
    model = load_face_model()

    with _helper_lock:
        helper = get_helper(upscale)
        helper.clean_all()
        helper.read_image(bgr)
        helper.get_face_landmarks_5(only_center_face=False, eye_dist_threshold=5)
        helper.align_warp_face()

        if not helper.cropped_faces:
            return background

        for cropped in helper.cropped_faces:
            # cropped is 512x512 BGR uint8; the model is scale 1.
            rgb = cropped[:, :, ::-1]
            restored = run_model(model, np.ascontiguousarray(rgb), tile=0)
            helper.add_restored_face(np.ascontiguousarray(restored[:, :, ::-1]))

        helper.get_inverse_affine(None)
        return helper.paste_faces_to_input_image(upsample_img=background)