feat: encpas the face restoration into a class
This commit is contained in:
@@ -1,9 +1,26 @@
|
||||
import os
|
||||
import time
|
||||
from PIL import Image, ImageOps, ImageSequence
|
||||
import shutil
|
||||
|
||||
import folder_paths
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ..filesystem import get_faceless_models
|
||||
from ..processors.face_restoration import FaceRestoration
|
||||
from ..vision import is_image
|
||||
|
||||
|
||||
class NodesFaceRestore:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
restoration_models = [os.path.basename(model) for model in get_faceless_models('face_restoration')]
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"restoration_model": (restoration_models,),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -13,7 +30,33 @@ class NodesFaceRestore:
|
||||
RETURN_NAMES = ("IMAGE",)
|
||||
FUNCTION = "restoreFace"
|
||||
|
||||
def restoreFace(self, images):
|
||||
print("restore face")
|
||||
return (images,)
|
||||
def restoreFace(self, images, restoration_model):
|
||||
face_restoration = FaceRestoration(restoration_model)
|
||||
|
||||
now = f"{int(time.time())}"
|
||||
output_path = os.path.join(folder_paths.get_temp_directory(), "faceless/restored_frames", now)
|
||||
if os.path.exists(output_path):
|
||||
shutil.rmtree(output_path)
|
||||
os.makedirs(output_path)
|
||||
|
||||
face_restoration.restore_images(images, output_path)
|
||||
|
||||
images = []
|
||||
for file in sorted(os.listdir(output_path)):
|
||||
file_path = os.path.join(output_path, file)
|
||||
if not is_image(file_path):
|
||||
continue
|
||||
img = Image.open(file_path)
|
||||
for i in ImageSequence.Iterator(img):
|
||||
i = ImageOps.exif_transpose(i)
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
image = i.convert("RGB")
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
images.append(image)
|
||||
if len(images) > 1:
|
||||
output_image = torch.cat(images, dim=0)
|
||||
else:
|
||||
output_image = images[0]
|
||||
return (output_image,)
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
from ..typing import FacelessVideo
|
||||
import os
|
||||
|
||||
from ..processors.face_enhancer import enhance_video
|
||||
from ..processors.face_restoration import FaceRestoration
|
||||
|
||||
from ..filesystem import get_faceless_models
|
||||
|
||||
class NodesVideoFaceRestore:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
restoration_models = [os.path.basename(model) for model in get_faceless_models('face_restoration')]
|
||||
return {
|
||||
"required": {
|
||||
"video": ("FACELESS_VIDEO",),
|
||||
"restoration_model": (restoration_models,),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -17,7 +21,8 @@ class NodesVideoFaceRestore:
|
||||
RETURN_NAMES = ("video",)
|
||||
FUNCTION = "restoreVideoFace"
|
||||
|
||||
def restoreVideoFace(self, video):
|
||||
def restoreVideoFace(self, video, restoration_model):
|
||||
frames_dir = video["frames_dir"]
|
||||
enhance_video(frames_dir)
|
||||
face_restoration = FaceRestoration(restoration_model)
|
||||
face_restoration.restore_video(frames_dir)
|
||||
return (video,)
|
||||
|
||||
@@ -1,223 +0,0 @@
|
||||
import os
|
||||
from queue import Queue
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
import threading
|
||||
from typing import List, Optional, Literal
|
||||
|
||||
import cv2
|
||||
import numpy
|
||||
import onnxruntime
|
||||
|
||||
from ..processors.face_analyser import get_many_faces
|
||||
from ..execution import apply_execution_provider_options
|
||||
from ..face_helper import warp_face_by_face_landmark_5, paste_back
|
||||
from ..face_masker import create_static_box_mask, create_occlusion_mask
|
||||
from ..vision import read_image, write_image
|
||||
from ..filesystem import resolve_relative_path
|
||||
from ..typing import VisionFrame, ModelSet, Any, OptionsWithModel
|
||||
|
||||
execution_thread_count = 4
|
||||
execution_queue_count = 1
|
||||
execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
|
||||
|
||||
model_template = "ffhq_512"
|
||||
model_size = (512, 512)
|
||||
model_path = (512, 512)
|
||||
face_mask_blur = 0.3
|
||||
face_mask_types = ['box']
|
||||
face_enhancer_blend = 80
|
||||
face_enhancer_model = 'gfpgan_1.4'
|
||||
|
||||
THREAD_LOCK : threading.Lock = threading.Lock()
|
||||
THREAD_SEMAPHORE : threading.Semaphore = threading.Semaphore()
|
||||
FRAME_PROCESSOR = None
|
||||
MODELS : ModelSet =\
|
||||
{
|
||||
'codeformer':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/codeformer.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/codeformer.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.2':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.2.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gfpgan_1.2.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.3':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.3.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gfpgan_1.3.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.4':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.4.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gfpgan_1.4.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gpen_bfr_256':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_256.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gpen_bfr_256.onnx'),
|
||||
'template': 'arcface_128_v2',
|
||||
'size': (256, 256)
|
||||
},
|
||||
'gpen_bfr_512':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_512.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gpen_bfr_512.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gpen_bfr_1024':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_1024.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gpen_bfr_1024.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (1024, 1024)
|
||||
},
|
||||
'gpen_bfr_2048':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_2048.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/gpen_bfr_2048.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (2048, 2048)
|
||||
},
|
||||
'restoreformer_plus_plus':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/restoreformer_plus_plus.onnx',
|
||||
'path': resolve_relative_path('../../../models/faceless/face_restoration/restoreformer_plus_plus.onnx'),
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
}
|
||||
}
|
||||
OPTIONS : Optional[OptionsWithModel] = None
|
||||
|
||||
def apply_enhance(crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
frame_processor = get_frame_processor()
|
||||
frame_processor_inputs = {}
|
||||
|
||||
for frame_processor_input in frame_processor.get_inputs():
|
||||
if frame_processor_input.name == 'input':
|
||||
frame_processor_inputs[frame_processor_input.name] = crop_vision_frame
|
||||
if frame_processor_input.name == 'weight':
|
||||
weight = numpy.array([ 1 ]).astype(numpy.double)
|
||||
frame_processor_inputs[frame_processor_input.name] = weight
|
||||
with THREAD_SEMAPHORE:
|
||||
crop_vision_frame = frame_processor.run(None, frame_processor_inputs)[0][0]
|
||||
return crop_vision_frame
|
||||
|
||||
def enhance_face(face, frame: VisionFrame) -> VisionFrame:
|
||||
crop_vision_frame, affine_matrix = warp_face_by_face_landmark_5(frame, face.landmarks.get('5/68'), model_template, model_size)
|
||||
box_mask = create_static_box_mask(crop_vision_frame.shape[:2][::-1], face_mask_blur, (0, 0, 0, 0))
|
||||
crop_mask_list =\
|
||||
[
|
||||
box_mask
|
||||
]
|
||||
|
||||
if 'occlusion' in face_mask_types:
|
||||
occlusion_mask = create_occlusion_mask(crop_vision_frame)
|
||||
crop_mask_list.append(occlusion_mask)
|
||||
crop_vision_frame = prepare_crop_frame(crop_vision_frame)
|
||||
crop_vision_frame = apply_enhance(crop_vision_frame)
|
||||
crop_vision_frame = normalize_crop_frame(crop_vision_frame)
|
||||
crop_mask = numpy.minimum.reduce(crop_mask_list).clip(0, 1)
|
||||
paste_vision_frame = paste_back(frame, crop_vision_frame, crop_mask, affine_matrix)
|
||||
temp_vision_frame = blend_frame(frame, paste_vision_frame)
|
||||
return temp_vision_frame
|
||||
|
||||
|
||||
def process_frame(frame: VisionFrame):
|
||||
# Support one face and many face mode
|
||||
faces = get_many_faces(frame)
|
||||
target_vision_frame = None
|
||||
for face in faces:
|
||||
target_vision_frame = enhance_face(face, frame)
|
||||
return target_vision_frame
|
||||
|
||||
def process_frames(target_frames_dir: str, queue_payloads: List[str]):
|
||||
count = len(queue_payloads)
|
||||
for index, frame_filename in enumerate(queue_payloads):
|
||||
print(f"progress: {index + 1}/{count}")
|
||||
frame_filepath = os.path.join(target_frames_dir, frame_filename)
|
||||
|
||||
target_vision_frame = read_image(frame_filepath)
|
||||
if target_vision_frame is None:
|
||||
raise Exception("invalid target image")
|
||||
output_vision_frame = process_frame(target_vision_frame)
|
||||
if output_vision_frame is None:
|
||||
continue
|
||||
# raise Exception("process frame failed")
|
||||
write_image(frame_filepath, output_vision_frame)
|
||||
|
||||
def enhance_video(frames_dir: str):
|
||||
frames_filenames = os.listdir(frames_dir)
|
||||
queue_payloads = sorted(frames_filenames)
|
||||
|
||||
with ThreadPoolExecutor(max_workers = execution_thread_count) as executor:
|
||||
futures = []
|
||||
queue : Queue[str] = create_queue(queue_payloads)
|
||||
queue_per_future = max(len(queue_payloads) // execution_thread_count * execution_queue_count, 1)
|
||||
while not queue.empty():
|
||||
future = executor.submit(process_frames, frames_dir, pick_queue(queue, queue_per_future))
|
||||
futures.append(future)
|
||||
for future_done in as_completed(futures):
|
||||
future_done.result()
|
||||
|
||||
def create_queue(queue_payloads : List[str]) -> Queue[str]:
|
||||
queue : Queue[str] = Queue()
|
||||
for queue_payload in queue_payloads:
|
||||
queue.put(queue_payload)
|
||||
return queue
|
||||
|
||||
def pick_queue(queue : Queue[str], queue_per_future : int) -> List[str]:
|
||||
queues = []
|
||||
for _ in range(queue_per_future):
|
||||
if not queue.empty():
|
||||
queues.append(queue.get())
|
||||
return queues
|
||||
|
||||
|
||||
def prepare_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
crop_vision_frame = crop_vision_frame[:, :, ::-1] / 255.0
|
||||
crop_vision_frame = (crop_vision_frame - 0.5) / 0.5
|
||||
crop_vision_frame = numpy.expand_dims(crop_vision_frame.transpose(2, 0, 1), axis = 0).astype(numpy.float32)
|
||||
return crop_vision_frame
|
||||
|
||||
def normalize_crop_frame(crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
crop_vision_frame = numpy.clip(crop_vision_frame, -1, 1)
|
||||
crop_vision_frame = (crop_vision_frame + 1) / 2
|
||||
crop_vision_frame = crop_vision_frame.transpose(1, 2, 0)
|
||||
crop_vision_frame = (crop_vision_frame * 255.0).round()
|
||||
crop_vision_frame = crop_vision_frame.astype(numpy.uint8)[:, :, ::-1]
|
||||
return crop_vision_frame
|
||||
|
||||
def blend_frame(temp_vision_frame : VisionFrame, paste_vision_frame : VisionFrame) -> VisionFrame:
|
||||
final_face_enhancer_blend = 1 - (face_enhancer_blend / 100)
|
||||
temp_vision_frame = cv2.addWeighted(temp_vision_frame, final_face_enhancer_blend, paste_vision_frame, 1 - final_face_enhancer_blend, 0)
|
||||
return temp_vision_frame
|
||||
|
||||
def get_options(key : Literal['model']) -> Any:
|
||||
global OPTIONS
|
||||
|
||||
if OPTIONS is None:
|
||||
OPTIONS =\
|
||||
{
|
||||
'model': MODELS[face_enhancer_model]
|
||||
}
|
||||
return OPTIONS.get(key)
|
||||
|
||||
def get_frame_processor() -> Any:
|
||||
global FRAME_PROCESSOR
|
||||
|
||||
with THREAD_LOCK:
|
||||
if FRAME_PROCESSOR is None:
|
||||
model_path = get_options('model').get('path')
|
||||
FRAME_PROCESSOR = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(execution_providers))
|
||||
return FRAME_PROCESSOR
|
||||
@@ -0,0 +1,225 @@
|
||||
import os
|
||||
from queue import Queue
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
import threading
|
||||
from typing import List
|
||||
|
||||
import cv2
|
||||
import numpy
|
||||
import onnxruntime
|
||||
|
||||
from ..processors.face_analyser import get_many_faces
|
||||
from ..execution import apply_execution_provider_options
|
||||
from ..face_helper import warp_face_by_face_landmark_5, paste_back
|
||||
from ..face_masker import create_static_box_mask, create_occlusion_mask
|
||||
from ..vision import read_image, write_image, tensor_to_vision_frame
|
||||
from ..typing import VisionFrame, ModelSet, Any
|
||||
from ..filesystem import get_faceless_model_path
|
||||
|
||||
THREAD_LOCK : threading.Lock = threading.Lock()
|
||||
THREAD_SEMAPHORE : threading.Semaphore = threading.Semaphore()
|
||||
|
||||
MODELS : ModelSet =\
|
||||
{
|
||||
'codeformer':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/codeformer.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.2':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.2.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.3':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.3.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gfpgan_1.4':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gfpgan_1.4.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gpen_bfr_256':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_256.onnx',
|
||||
'template': 'arcface_128_v2',
|
||||
'size': (256, 256)
|
||||
},
|
||||
'gpen_bfr_512':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_512.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
},
|
||||
'gpen_bfr_1024':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_1024.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (1024, 1024)
|
||||
},
|
||||
'gpen_bfr_2048':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/gpen_bfr_2048.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (2048, 2048)
|
||||
},
|
||||
'restoreformer_plus_plus':
|
||||
{
|
||||
'url': 'https://github.com/facefusion/facefusion-assets/releases/download/models/restoreformer_plus_plus.onnx',
|
||||
'template': 'ffhq_512',
|
||||
'size': (512, 512)
|
||||
}
|
||||
}
|
||||
|
||||
class FaceRestoration:
|
||||
|
||||
def __init__(self, model_name: str) -> None:
|
||||
self._model_name = model_name
|
||||
|
||||
self._execution_thread_count = 4
|
||||
self._execution_queue_count = 1
|
||||
self._execution_providers = ['CoreMLExecutionProvider', 'CPUExecutionProvider']
|
||||
|
||||
self._frame_processor = None
|
||||
|
||||
self._face_mask_blur = 0.3
|
||||
self._face_mask_types = ['box']
|
||||
self._face_enhancer_blend = 80
|
||||
self._face_enhancer_model = 'gfpgan_1.4'
|
||||
|
||||
def restore_images(self, images, output_path: str):
|
||||
for (index, image) in enumerate(images):
|
||||
filename = f"{index + 1}".ljust(4, "0") + ".png"
|
||||
output_filepath = os.path.join(output_path, filename)
|
||||
|
||||
target_vision_frame = tensor_to_vision_frame(image)
|
||||
if target_vision_frame is None:
|
||||
raise Exception("invalid target image")
|
||||
output_vision_frame = self._process_frame(target_vision_frame)
|
||||
if output_vision_frame is None:
|
||||
continue
|
||||
# raise Exception("process frame failed")
|
||||
write_image(output_filepath, output_vision_frame)
|
||||
|
||||
def restore_video(self, frames_dir: str):
|
||||
frames_filenames = os.listdir(frames_dir)
|
||||
queue_payloads = sorted(frames_filenames)
|
||||
|
||||
with ThreadPoolExecutor(max_workers = self._execution_thread_count) as executor:
|
||||
futures = []
|
||||
queue : Queue[str] = self._create_queue(queue_payloads)
|
||||
queue_per_future = max(len(queue_payloads) // self._execution_thread_count * self._execution_queue_count, 1)
|
||||
while not queue.empty():
|
||||
future = executor.submit(self._process_frames, frames_dir, self._pick_queue(queue, queue_per_future))
|
||||
futures.append(future)
|
||||
for future_done in as_completed(futures):
|
||||
future_done.result()
|
||||
|
||||
def _create_queue(self, queue_payloads : List[str]) -> Queue[str]:
|
||||
queue : Queue[str] = Queue()
|
||||
for queue_payload in queue_payloads:
|
||||
queue.put(queue_payload)
|
||||
return queue
|
||||
|
||||
def _pick_queue(self, queue : Queue[str], queue_per_future : int) -> List[str]:
|
||||
queues = []
|
||||
for _ in range(queue_per_future):
|
||||
if not queue.empty():
|
||||
queues.append(queue.get())
|
||||
return queues
|
||||
|
||||
def _process_frames(self, target_frames_dir: str, queue_payloads: List[str]):
|
||||
count = len(queue_payloads)
|
||||
for index, frame_filename in enumerate(queue_payloads):
|
||||
print(f"progress: {index + 1}/{count}")
|
||||
frame_filepath = os.path.join(target_frames_dir, frame_filename)
|
||||
|
||||
target_vision_frame = read_image(frame_filepath)
|
||||
if target_vision_frame is None:
|
||||
raise Exception("invalid target image")
|
||||
output_vision_frame = self._process_frame(target_vision_frame)
|
||||
if output_vision_frame is None:
|
||||
continue
|
||||
# raise Exception("process frame failed")
|
||||
write_image(frame_filepath, output_vision_frame)
|
||||
|
||||
def _process_frame(self, frame: VisionFrame):
|
||||
# Support one face and many face mode
|
||||
faces = get_many_faces(frame)
|
||||
target_vision_frame = None
|
||||
for face in faces:
|
||||
target_vision_frame = self._enhance_face(face, frame)
|
||||
return target_vision_frame
|
||||
|
||||
def _enhance_face(self, face, frame: VisionFrame) -> VisionFrame:
|
||||
model_template = self._get_model_options().get('template')
|
||||
model_size = self._get_model_options().get('size')
|
||||
crop_vision_frame, affine_matrix = warp_face_by_face_landmark_5(frame, face.landmarks.get('5/68'), model_template, model_size)
|
||||
box_mask = create_static_box_mask(crop_vision_frame.shape[:2][::-1], self._face_mask_blur, (0, 0, 0, 0))
|
||||
crop_mask_list =\
|
||||
[
|
||||
box_mask
|
||||
]
|
||||
|
||||
if 'occlusion' in self._face_mask_types:
|
||||
occlusion_mask = create_occlusion_mask(crop_vision_frame)
|
||||
crop_mask_list.append(occlusion_mask)
|
||||
crop_vision_frame = self._prepare_crop_frame(crop_vision_frame)
|
||||
crop_vision_frame = self._apply_enhance(crop_vision_frame)
|
||||
crop_vision_frame = self._normalize_crop_frame(crop_vision_frame)
|
||||
crop_mask = numpy.minimum.reduce(crop_mask_list).clip(0, 1)
|
||||
paste_vision_frame = paste_back(frame, crop_vision_frame, crop_mask, affine_matrix)
|
||||
temp_vision_frame = self._blend_frame(frame, paste_vision_frame)
|
||||
return temp_vision_frame
|
||||
|
||||
def _prepare_crop_frame(self, crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
crop_vision_frame = crop_vision_frame[:, :, ::-1] / 255.0
|
||||
crop_vision_frame = (crop_vision_frame - 0.5) / 0.5
|
||||
crop_vision_frame = numpy.expand_dims(crop_vision_frame.transpose(2, 0, 1), axis = 0).astype(numpy.float32)
|
||||
return crop_vision_frame
|
||||
|
||||
def _apply_enhance(self, crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
frame_processor = self._get_frame_processor()
|
||||
frame_processor_inputs = {}
|
||||
|
||||
for frame_processor_input in frame_processor.get_inputs():
|
||||
if frame_processor_input.name == 'input':
|
||||
frame_processor_inputs[frame_processor_input.name] = crop_vision_frame
|
||||
if frame_processor_input.name == 'weight':
|
||||
weight = numpy.array([ 1 ]).astype(numpy.double)
|
||||
frame_processor_inputs[frame_processor_input.name] = weight
|
||||
with THREAD_SEMAPHORE:
|
||||
crop_vision_frame = frame_processor.run(None, frame_processor_inputs)[0][0]
|
||||
return crop_vision_frame
|
||||
|
||||
def _normalize_crop_frame(self, crop_vision_frame : VisionFrame) -> VisionFrame:
|
||||
crop_vision_frame = numpy.clip(crop_vision_frame, -1, 1)
|
||||
crop_vision_frame = (crop_vision_frame + 1) / 2
|
||||
crop_vision_frame = crop_vision_frame.transpose(1, 2, 0)
|
||||
crop_vision_frame = (crop_vision_frame * 255.0).round()
|
||||
crop_vision_frame = crop_vision_frame.astype(numpy.uint8)[:, :, ::-1]
|
||||
return crop_vision_frame
|
||||
|
||||
def _get_frame_processor(self) -> Any:
|
||||
global FRAME_PROCESSOR
|
||||
|
||||
with THREAD_LOCK:
|
||||
if self._frame_processor is None:
|
||||
model_path = get_faceless_model_path('face_restoration', self._model_name)
|
||||
self._frame_processor = onnxruntime.InferenceSession(model_path, providers = apply_execution_provider_options(self._execution_providers))
|
||||
return self._frame_processor
|
||||
|
||||
def _blend_frame(self, temp_vision_frame : VisionFrame, paste_vision_frame : VisionFrame) -> VisionFrame:
|
||||
final_face_enhancer_blend = 1 - (self._face_enhancer_blend / 100)
|
||||
temp_vision_frame = cv2.addWeighted(temp_vision_frame, final_face_enhancer_blend, paste_vision_frame, 1 - final_face_enhancer_blend, 0)
|
||||
return temp_vision_frame
|
||||
|
||||
def _get_model_options(self) -> Any:
|
||||
names = os.path.splitext(self._model_name)
|
||||
return MODELS.get(names[0])
|
||||
+1
-1
@@ -72,7 +72,7 @@ FaceStore = TypedDict('FaceStore',
|
||||
})
|
||||
|
||||
|
||||
ModelType = Literal['face_swapper', 'face_detector', 'face_recognizer', 'face_landmarker']
|
||||
ModelType = Literal['face_swapper', 'face_restoration', 'face_detector', 'face_recognizer', 'face_landmarker']
|
||||
FaceDetectorModel = Literal['many', 'retinaface', 'scrfd', 'yoloface', 'yunet']
|
||||
FaceRecognizerModel = Literal['arcface_blendswap', 'arcface_inswapper', 'arcface_simswap', 'arcface_uniface']
|
||||
|
||||
|
||||
Reference in New Issue
Block a user