From d5c9d3c4268fbbed006423ef8ae2e003ea71281b Mon Sep 17 00:00:00 2001 From: Jeffrey Wu Date: Mon, 29 Jul 2024 11:59:56 +0800 Subject: [PATCH] feat: encpas the face restoration into a class --- faceless/nodes/nodes_face_restore.py | 49 ++++- faceless/nodes/nodes_video_face_restore.py | 13 +- faceless/processors/face_enhancer.py | 223 -------------------- faceless/processors/face_restoration.py | 225 +++++++++++++++++++++ faceless/typing.py | 2 +- 5 files changed, 281 insertions(+), 231 deletions(-) delete mode 100644 faceless/processors/face_enhancer.py create mode 100644 faceless/processors/face_restoration.py diff --git a/faceless/nodes/nodes_face_restore.py b/faceless/nodes/nodes_face_restore.py index 80fd75d..726090b 100644 --- a/faceless/nodes/nodes_face_restore.py +++ b/faceless/nodes/nodes_face_restore.py @@ -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,) diff --git a/faceless/nodes/nodes_video_face_restore.py b/faceless/nodes/nodes_video_face_restore.py index e0f7fb7..4d77f98 100644 --- a/faceless/nodes/nodes_video_face_restore.py +++ b/faceless/nodes/nodes_video_face_restore.py @@ -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,) diff --git a/faceless/processors/face_enhancer.py b/faceless/processors/face_enhancer.py deleted file mode 100644 index 75e923a..0000000 --- a/faceless/processors/face_enhancer.py +++ /dev/null @@ -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 diff --git a/faceless/processors/face_restoration.py b/faceless/processors/face_restoration.py new file mode 100644 index 0000000..11b586d --- /dev/null +++ b/faceless/processors/face_restoration.py @@ -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]) diff --git a/faceless/typing.py b/faceless/typing.py index 21fcf3c..56e62ac 100644 --- a/faceless/typing.py +++ b/faceless/typing.py @@ -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']