feat: encpas the face restoration into a class

This commit is contained in:
Jeffrey Wu
2024-07-29 11:59:56 +08:00
parent 92c1cda36e
commit d5c9d3c426
5 changed files with 281 additions and 231 deletions
+46 -3
View File
@@ -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,)
+9 -4
View File
@@ -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,)
-223
View File
@@ -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
+225
View File
@@ -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
View File
@@ -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']