From 3bb861b6081bc0fd3223baa2a0b4dd5a66baa0fe Mon Sep 17 00:00:00 2001 From: Jeffrey Wu Date: Mon, 13 May 2024 17:44:14 +0800 Subject: [PATCH] feat: support remove background --- faceless/ffmpeg.py | 6 +- faceless/image_helper.py | 10 + faceless/nodes/globals.py | 7 + faceless/nodes/nodes_load_frames.py | 2 +- faceless/nodes/nodes_load_video.py | 13 +- faceless/nodes/nodes_load_video_url.py | 25 +- faceless/nodes/nodes_remove_background.py | 104 ++++ faceless/nodes/nodes_save_video.py | 18 +- faceless/nodes/nodes_video_face_swap.py | 14 +- .../nodes/nodes_video_remove_background.py | 44 ++ faceless/processors/briarmbg.py | 457 ++++++++++++++++++ faceless/processors/face_swapper.py | 13 +- faceless/typing.py | 4 + 13 files changed, 668 insertions(+), 49 deletions(-) create mode 100644 faceless/image_helper.py create mode 100644 faceless/nodes/nodes_remove_background.py create mode 100644 faceless/nodes/nodes_video_remove_background.py create mode 100644 faceless/processors/briarmbg.py diff --git a/faceless/ffmpeg.py b/faceless/ffmpeg.py index 12c93bb..657adfd 100644 --- a/faceless/ffmpeg.py +++ b/faceless/ffmpeg.py @@ -43,9 +43,9 @@ def extract_frames(video_path: str, frames_path: str, video_resolution : Resolut commands.extend([ '-vsync', '0', temp_frames_pattern ]) return run_ffmpeg(commands) -def merge_video(target_path: str, output_path: str, video_resolution: Resolution, video_fps: Fps, output_video_encoder: OutputVideoEncoder = 'libx264', output_video_quality: int = 80, output_video_preset: OutputVideoPreset = 'veryfast', frame_format: FrameFormat = 'png') -> bool: - temp_video_fps = restrict_video_fps(target_path, video_fps) - temp_frames_pattern = get_temp_frames_pattern(target_path, '%04d', frame_format) +def merge_video(video_path: str, frames_dir: str, output_path: str, video_resolution: Resolution, video_fps: Fps, output_video_encoder: OutputVideoEncoder = 'libx264', output_video_quality: int = 80, output_video_preset: OutputVideoPreset = 'veryfast', frame_format: FrameFormat = 'png') -> bool: + temp_video_fps = restrict_video_fps(video_path, video_fps) + temp_frames_pattern = get_temp_frames_pattern(frames_dir, '%04d', frame_format) commands = [ '-hwaccel', 'auto', '-s', pack_resolution(video_resolution), '-r', str(temp_video_fps), '-i', temp_frames_pattern, '-c:v', output_video_encoder ] if output_video_encoder in [ 'libx264', 'libx265' ]: diff --git a/faceless/image_helper.py b/faceless/image_helper.py new file mode 100644 index 0000000..d077f02 --- /dev/null +++ b/faceless/image_helper.py @@ -0,0 +1,10 @@ +from PIL import Image + +import torch +import numpy as np + +def tensor_to_pil(image): + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + +def pil_to_tensor(image): + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) diff --git a/faceless/nodes/globals.py b/faceless/nodes/globals.py index 0b67eb8..2a83dec 100644 --- a/faceless/nodes/globals.py +++ b/faceless/nodes/globals.py @@ -7,8 +7,10 @@ from .nodes_upload_video import NodesUploadVideo from .nodes_face_swap import NodesFaceSwap from .nodes_face_restore import NodesFaceRestore +from .nodes_remove_background import NodesRemoveBackground from .nodes_video_face_swap import NodesVideoFaceSwap from .nodes_video_face_restore import NodesVideoFaceRestore +from .nodes_video_remove_background import NodesVideoRemoveBackground NODE_CLASS_MAPPINGS = { "FacelessLoadVideo": NodesLoadVideo, @@ -20,8 +22,11 @@ NODE_CLASS_MAPPINGS = { "FacelessFaceSwap": NodesFaceSwap, "FacelessFaceRestore": NodesFaceRestore, + "FacelessRemoveBackground": NodesRemoveBackground, + "FacelessVideoFaceSwap": NodesVideoFaceSwap, "FacelessVideoFaceRestore": NodesVideoFaceRestore, + "FacelessVideoRemoveBackground": NodesVideoRemoveBackground, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -36,4 +41,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FacelessFaceRestore": "Face Restore", "FacelessVideoFaceSwap": "Face Swap (Video)", "FacelessVideoFaceRestore": "Face Restore (Video)", + "FacelessRemoveBackground": "Remove Background", + "FacelessVideoRemoveBackground": "Remove Background (Video)", } diff --git a/faceless/nodes/nodes_load_frames.py b/faceless/nodes/nodes_load_frames.py index a524c42..1fdbd98 100644 --- a/faceless/nodes/nodes_load_frames.py +++ b/faceless/nodes/nodes_load_frames.py @@ -21,7 +21,7 @@ class NodesLoadFrames: FUNCTION = "load_frames" def load_frames(self, video: FacelessVideo): - frames_path = video['output_path'] + frames_path = video["frames_dir"] images = [] for file in sorted(os.listdir(frames_path)): diff --git a/faceless/nodes/nodes_load_video.py b/faceless/nodes/nodes_load_video.py index d8eeef4..bc0780d 100644 --- a/faceless/nodes/nodes_load_video.py +++ b/faceless/nodes/nodes_load_video.py @@ -75,11 +75,12 @@ class NodesLoadVideo: raise Exception("Failed to extract frames") faceless_video: FacelessVideo = { - 'video_path': video_path, - 'output_path': frames_path, - 'resolution': video_resolution, - 'fps': video_fps, - 'trim_frame_start': final_trim_frame_start, - 'trim_frame_end': final_trim_frame_end, + "video_path": video_path, + "frames_dir": frames_path, + "output_path": "", + "resolution": video_resolution, + "fps": video_fps, + "trim_frame_start": final_trim_frame_start, + "trim_frame_end": final_trim_frame_end, } return (faceless_video,) diff --git a/faceless/nodes/nodes_load_video_url.py b/faceless/nodes/nodes_load_video_url.py index 136e138..116b6b1 100644 --- a/faceless/nodes/nodes_load_video_url.py +++ b/faceless/nodes/nodes_load_video_url.py @@ -63,13 +63,13 @@ class NodesLoadVideoUrl: # Save video video_name, _ = os.path.splitext(os.path.basename(video_filepath)) - frames_path = os.path.join(folder_paths.get_temp_directory(), "faceless/frames", video_name) - print("frames path: " + frames_path) + frames_dir = os.path.join(folder_paths.get_temp_directory(), "faceless/frames", video_name) + print("frames path: " + frames_dir) # Remove all cached frames - if os.path.exists(frames_path): - shutil.rmtree(frames_path) - os.makedirs(frames_path) + if os.path.exists(frames_dir): + shutil.rmtree(frames_dir) + os.makedirs(frames_dir) video_resolution = detect_video_resolution(video_filepath) video_fps = detect_video_fps(video_filepath) @@ -85,16 +85,17 @@ class NodesLoadVideoUrl: else: final_trim_frame_end = trim_frame_end - if not extract_frames(video_filepath, frames_path, video_resolution, video_fps, final_trim_frame_start, final_trim_frame_end): + if not extract_frames(video_filepath, frames_dir, video_resolution, video_fps, final_trim_frame_start, final_trim_frame_end): raise Exception("Failed to extract frames") faceless_video: FacelessVideo = { - 'video_path': video_filepath, - 'output_path': frames_path, - 'resolution': video_resolution, - 'fps': video_fps, - 'trim_frame_start': final_trim_frame_start, - 'trim_frame_end': final_trim_frame_end, + "video_path": video_filepath, + "frames_dir": frames_dir, + "output_path": "", + "resolution": video_resolution, + "fps": video_fps, + "trim_frame_start": final_trim_frame_start, + "trim_frame_end": final_trim_frame_end, } return (faceless_video,) diff --git a/faceless/nodes/nodes_remove_background.py b/faceless/nodes/nodes_remove_background.py new file mode 100644 index 0000000..ab32f4d --- /dev/null +++ b/faceless/nodes/nodes_remove_background.py @@ -0,0 +1,104 @@ +import os +from PIL import Image + +import torch +import torch.nn.functional as F +import numpy as np +from torchvision.transforms.functional import normalize + +from folder_paths import models_dir + +from ..processors.briarmbg import BriaRMBG +from ..image_helper import tensor_to_pil, pil_to_tensor + +class NodesRemoveBackground: + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "images": ("IMAGE",), + }, + } + + CATEGORY = "faceless" + RETURN_TYPES = ("IMAGE", "MASK",) + FUNCTION = "remove_images_background" + + @classmethod + def VALIDATE_INPUTS(cls, images): + if not os.path.exists(os.path.join(models_dir, "faceless/rmbg.pth")): + return False + return True + + def remove_images_background(self, images): + self.load_model() + + processed_images = [] + processed_masks = [] + for image in images: + orig_image = tensor_to_pil(image) + + new_im, pil_im = self.remove_background(orig_image) + + new_im_tensor = pil_to_tensor(new_im) + pil_im_tensor = pil_to_tensor(pil_im) + + processed_images.append(new_im_tensor) + processed_masks.append(pil_im_tensor) + + new_ims = torch.cat(processed_images, dim=0) + new_masks = torch.cat(processed_masks, dim=0) + return (new_ims, new_masks) + + def remove_background(self, orig_image): + w, h = orig_image.size + model_input_size = [1024,1024] + image = self._preprocess_image(np.array(orig_image), model_input_size) + + if torch.cuda.is_available(): + image = image.to("cuda") + elif torch.backends.mps.is_available(): + image = image.to("mps") + + result = self.rmbg(image) + + result_image = self._postprocess_image(result[0][0], [h, w]) + + pil_im = Image.fromarray(result_image) + no_bg_image = Image.new("RGBA", pil_im.size, (0,0,0,0)) + no_bg_image.paste(orig_image, mask=pil_im) + return (no_bg_image, pil_im) + + def load_model(self): + rmbg = BriaRMBG() + if torch.cuda.is_available(): + device = "cuda" + elif torch.backends.mps.is_available(): + device = "mps" + else: + device = "cpu" + model_path = os.path.join(models_dir, "faceless/rmbg.pth") + rmbg.load_state_dict(torch.load(model_path, map_location=device)) + rmbg.to(device) + rmbg.eval() + self.rmbg = rmbg + + def _preprocess_image(self, im: np.ndarray, model_input_size: list) -> torch.Tensor: + if len(im.shape) < 3: + im = im[:, :, np.newaxis] + # orig_im_size=im.shape[0:2] + im_tensor = torch.tensor(im, dtype=torch.float32).permute(2,0,1) + im_tensor = F.interpolate(torch.unsqueeze(im_tensor,0), size=model_input_size, mode='bilinear') + image = torch.divide(im_tensor,255.0) + image = normalize(image,[0.5,0.5,0.5],[1.0,1.0,1.0]) + return image + + def _postprocess_image(self, result: torch.Tensor, im_size: list)-> np.ndarray: + result = torch.squeeze(F.interpolate(result, size=im_size, mode='bilinear') ,0) + ma = torch.max(result) + mi = torch.min(result) + result = (result-mi)/(ma-mi) + im_array = (result*255).permute(1,2,0).cpu().data.numpy().astype(np.uint8) + im_array = np.squeeze(im_array) + return im_array diff --git a/faceless/nodes/nodes_save_video.py b/faceless/nodes/nodes_save_video.py index 6bd048b..29bdbb8 100644 --- a/faceless/nodes/nodes_save_video.py +++ b/faceless/nodes/nodes_save_video.py @@ -24,29 +24,31 @@ class NodesSaveVideo: def save_video(self, video: FacelessVideo): video_path = video.get("video_path") - frames_path = video.get("output_path") + frames_dir = video.get("frames_dir") fps = video.get("fps") trim_frame_start = video.get("trim_frame_start") trim_frame_end = video.get("trim_frame_end") - output_dir = os.path.join(folder_paths.get_output_directory(), "faceless") - if not os.path.exists(output_dir): - os.makedirs(output_dir) - - resolution = video.get("resolution") - fps = video.get("fps") + resolution = video["resolution"] + fps = video["fps"] output_temp_path = os.path.join(folder_paths.get_temp_directory(), "faceless/output", os.path.basename(video_path)) if not os.path.exists(os.path.dirname(output_temp_path)): os.makedirs(os.path.dirname(output_temp_path)) - if not merge_video(frames_path, output_temp_path, resolution, fps): + # Merge frames + if not merge_video(video_path, frames_dir, output_temp_path, resolution, fps): raise Exception("Failed to merge video") + # Restore audio + output_dir = os.path.join(folder_paths.get_output_directory(), "faceless") + if not os.path.exists(output_dir): + os.makedirs(output_dir) now = int(time.time()) output_path = os.path.join(output_dir, f"{now}_" + os.path.basename(video_path)) if not restore_audio(output_temp_path, video_path, output_path, fps, trim_frame_start, trim_frame_end): raise Exception("Failed to restore audio") + video["output_path"] = output_path return () diff --git a/faceless/nodes/nodes_video_face_swap.py b/faceless/nodes/nodes_video_face_swap.py index 93b2490..39048ad 100644 --- a/faceless/nodes/nodes_video_face_swap.py +++ b/faceless/nodes/nodes_video_face_swap.py @@ -31,20 +31,10 @@ class NodesVideoFaceSwap: FUNCTION = "swap_video_face" def swap_video_face(self, source_image, target_video: FacelessVideo, swapper_model, detector_model, recognizer_model): - video_path = target_video.get("video_path") - video_name, _ = os.path.splitext(os.path.basename(video_path)) - - frames_path = os.path.join(folder_paths.get_temp_directory(), "faceless/frames", video_name) - output_path = os.path.join(folder_paths.get_temp_directory(), "faceless/swapped_frames", video_name) - - if os.path.exists(output_path): - shutil.rmtree(output_path) - os.makedirs(output_path) + frames_path = target_video["frames_dir"] # TODO Check if has face on source image # Fetch source image or change process_frames argument. - swap_video(source_image[0], frames_path, output_path) - - target_video['output_path'] = output_path + swap_video(source_image[0], frames_path) return (target_video,) diff --git a/faceless/nodes/nodes_video_remove_background.py b/faceless/nodes/nodes_video_remove_background.py new file mode 100644 index 0000000..d10a049 --- /dev/null +++ b/faceless/nodes/nodes_video_remove_background.py @@ -0,0 +1,44 @@ +import os +from PIL import Image + +from ..vision import is_image +from ..typing import FacelessVideo +from .nodes_remove_background import NodesRemoveBackground +import torch.multiprocessing as mp + +class NodesVideoRemoveBackground(NodesRemoveBackground): + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "video": ("FACELESS_VIDEO",), + }, + } + + CATEGORY = "faceless" + RETURN_TYPES = () + RETURN_TYPES = ("FACELESS_VIDEO",) + RETURN_NAMES = ("video",) + FUNCTION = "remove_video_background" + + @classmethod + def VALIDATE_INPUTS(cls, video): + return super().VALIDATE_INPUTS(()) + + def remove_video_background(self, video: FacelessVideo): + frames_dir = video["frames_dir"] + + self.load_model() + + # TODO Improve batch process performance + frame_filenames = sorted(os.listdir(frames_dir)) + for frame_filename in frame_filenames: + file_path = os.path.join(frames_dir, frame_filename) + if not is_image(file_path): + continue + + img = Image.open(file_path) + new_im, _ = self.remove_background(img) + new_im.save(file_path) + return (video,) diff --git a/faceless/processors/briarmbg.py b/faceless/processors/briarmbg.py new file mode 100644 index 0000000..59d1eef --- /dev/null +++ b/faceless/processors/briarmbg.py @@ -0,0 +1,457 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class REBNCONV(nn.Module): + def __init__(self, in_ch=3, out_ch=3, dirate=1, stride=1): + super(REBNCONV, self).__init__() + + self.conv_s1 = nn.Conv2d( + in_ch, out_ch, 3, padding=1 * dirate, dilation=1 * dirate, stride=stride + ) + self.bn_s1 = nn.BatchNorm2d(out_ch) + self.relu_s1 = nn.ReLU(inplace=True) + + def forward(self, x): + hx = x + xout = self.relu_s1(self.bn_s1(self.conv_s1(hx))) + + return xout + + +def _upsample_like(src, tar): + src = F.interpolate(src, size=tar.shape[2:], mode="bilinear") + return src + + +### RSU-7 ### +class RSU7(nn.Module): + def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512): + super(RSU7, self).__init__() + + self.in_ch = in_ch + self.mid_ch = mid_ch + self.out_ch = out_ch + + self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) ## 1 -> 1/2 + + self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) + self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool5 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=1) + + self.rebnconv7 = REBNCONV(mid_ch, mid_ch, dirate=2) + + self.rebnconv6d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv5d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1) + + def forward(self, x): + b, c, h, w = x.shape + + hx = x + hxin = self.rebnconvin(hx) + + hx1 = self.rebnconv1(hxin) + hx = self.pool1(hx1) + + hx2 = self.rebnconv2(hx) + hx = self.pool2(hx2) + + hx3 = self.rebnconv3(hx) + hx = self.pool3(hx3) + + hx4 = self.rebnconv4(hx) + hx = self.pool4(hx4) + + hx5 = self.rebnconv5(hx) + hx = self.pool5(hx5) + + hx6 = self.rebnconv6(hx) + + hx7 = self.rebnconv7(hx6) + + hx6d = self.rebnconv6d(torch.cat((hx7, hx6), 1)) + hx6dup = _upsample_like(hx6d, hx5) + + hx5d = self.rebnconv5d(torch.cat((hx6dup, hx5), 1)) + hx5dup = _upsample_like(hx5d, hx4) + + hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1)) + hx4dup = _upsample_like(hx4d, hx3) + + hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) + hx3dup = _upsample_like(hx3d, hx2) + + hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) + hx2dup = _upsample_like(hx2d, hx1) + + hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) + + return hx1d + hxin + + +### RSU-6 ### +class RSU6(nn.Module): + def __init__(self, in_ch=3, mid_ch=12, out_ch=3): + super(RSU6, self).__init__() + + self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) + + self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) + self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1) + + self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=2) + + self.rebnconv5d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1) + + def forward(self, x): + hx = x + + hxin = self.rebnconvin(hx) + + hx1 = self.rebnconv1(hxin) + hx = self.pool1(hx1) + + hx2 = self.rebnconv2(hx) + hx = self.pool2(hx2) + + hx3 = self.rebnconv3(hx) + hx = self.pool3(hx3) + + hx4 = self.rebnconv4(hx) + hx = self.pool4(hx4) + + hx5 = self.rebnconv5(hx) + + hx6 = self.rebnconv6(hx5) + + hx5d = self.rebnconv5d(torch.cat((hx6, hx5), 1)) + hx5dup = _upsample_like(hx5d, hx4) + + hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1)) + hx4dup = _upsample_like(hx4d, hx3) + + hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) + hx3dup = _upsample_like(hx3d, hx2) + + hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) + hx2dup = _upsample_like(hx2d, hx1) + + hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) + + return hx1d + hxin + + +### RSU-5 ### +class RSU5(nn.Module): + def __init__(self, in_ch=3, mid_ch=12, out_ch=3): + super(RSU5, self).__init__() + + self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) + + self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) + self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) + + self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=2) + + self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1) + + def forward(self, x): + hx = x + + hxin = self.rebnconvin(hx) + + hx1 = self.rebnconv1(hxin) + hx = self.pool1(hx1) + + hx2 = self.rebnconv2(hx) + hx = self.pool2(hx2) + + hx3 = self.rebnconv3(hx) + hx = self.pool3(hx3) + + hx4 = self.rebnconv4(hx) + + hx5 = self.rebnconv5(hx4) + + hx4d = self.rebnconv4d(torch.cat((hx5, hx4), 1)) + hx4dup = _upsample_like(hx4d, hx3) + + hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) + hx3dup = _upsample_like(hx3d, hx2) + + hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) + hx2dup = _upsample_like(hx2d, hx1) + + hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) + + return hx1d + hxin + + +### RSU-4 ### +class RSU4(nn.Module): + def __init__(self, in_ch=3, mid_ch=12, out_ch=3): + super(RSU4, self).__init__() + + self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) + + self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) + self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) + self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) + + self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=2) + + self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1) + self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1) + + def forward(self, x): + hx = x + + hxin = self.rebnconvin(hx) + + hx1 = self.rebnconv1(hxin) + hx = self.pool1(hx1) + + hx2 = self.rebnconv2(hx) + hx = self.pool2(hx2) + + hx3 = self.rebnconv3(hx) + + hx4 = self.rebnconv4(hx3) + + hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1)) + hx3dup = _upsample_like(hx3d, hx2) + + hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) + hx2dup = _upsample_like(hx2d, hx1) + + hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) + + return hx1d + hxin + + +### RSU-4F ### +class RSU4F(nn.Module): + def __init__(self, in_ch=3, mid_ch=12, out_ch=3): + super(RSU4F, self).__init__() + + self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) + + self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) + self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=2) + self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=4) + + self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=8) + + self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=4) + self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=2) + self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1) + + def forward(self, x): + hx = x + + hxin = self.rebnconvin(hx) + + hx1 = self.rebnconv1(hxin) + hx2 = self.rebnconv2(hx1) + hx3 = self.rebnconv3(hx2) + + hx4 = self.rebnconv4(hx3) + + hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1)) + hx2d = self.rebnconv2d(torch.cat((hx3d, hx2), 1)) + hx1d = self.rebnconv1d(torch.cat((hx2d, hx1), 1)) + + return hx1d + hxin + + +class myrebnconv(nn.Module): + def __init__( + self, + in_ch=3, + out_ch=1, + kernel_size=3, + stride=1, + padding=1, + dilation=1, + groups=1, + ): + super(myrebnconv, self).__init__() + + self.conv = nn.Conv2d( + in_ch, + out_ch, + kernel_size=kernel_size, + stride=stride, + padding=padding, + dilation=dilation, + groups=groups, + ) + self.bn = nn.BatchNorm2d(out_ch) + self.rl = nn.ReLU(inplace=True) + + def forward(self, x): + return self.rl(self.bn(self.conv(x))) + + +class BriaRMBG(nn.Module): + def __init__(self, config: dict = {"in_ch": 3, "out_ch": 1}): + super(BriaRMBG, self).__init__() + in_ch = config["in_ch"] + out_ch = config["out_ch"] + self.conv_in = nn.Conv2d(in_ch, 64, 3, stride=2, padding=1) + self.pool_in = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage1 = RSU7(64, 32, 64) + self.pool12 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage2 = RSU6(64, 32, 128) + self.pool23 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage3 = RSU5(128, 64, 256) + self.pool34 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage4 = RSU4(256, 128, 512) + self.pool45 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage5 = RSU4F(512, 256, 512) + self.pool56 = nn.MaxPool2d(2, stride=2, ceil_mode=True) + + self.stage6 = RSU4F(512, 256, 512) + + # decoder + self.stage5d = RSU4F(1024, 256, 512) + self.stage4d = RSU4(1024, 128, 256) + self.stage3d = RSU5(512, 64, 128) + self.stage2d = RSU6(256, 32, 64) + self.stage1d = RSU7(128, 16, 64) + + self.side1 = nn.Conv2d(64, out_ch, 3, padding=1) + self.side2 = nn.Conv2d(64, out_ch, 3, padding=1) + self.side3 = nn.Conv2d(128, out_ch, 3, padding=1) + self.side4 = nn.Conv2d(256, out_ch, 3, padding=1) + self.side5 = nn.Conv2d(512, out_ch, 3, padding=1) + self.side6 = nn.Conv2d(512, out_ch, 3, padding=1) + + # self.outconv = nn.Conv2d(6*out_ch,out_ch,1) + + def forward(self, x): + hx = x + + hxin = self.conv_in(hx) + # hx = self.pool_in(hxin) + + # stage 1 + hx1 = self.stage1(hxin) + hx = self.pool12(hx1) + + # stage 2 + hx2 = self.stage2(hx) + hx = self.pool23(hx2) + + # stage 3 + hx3 = self.stage3(hx) + hx = self.pool34(hx3) + + # stage 4 + hx4 = self.stage4(hx) + hx = self.pool45(hx4) + + # stage 5 + hx5 = self.stage5(hx) + hx = self.pool56(hx5) + + # stage 6 + hx6 = self.stage6(hx) + hx6up = _upsample_like(hx6, hx5) + + # -------------------- decoder -------------------- + hx5d = self.stage5d(torch.cat((hx6up, hx5), 1)) + hx5dup = _upsample_like(hx5d, hx4) + + hx4d = self.stage4d(torch.cat((hx5dup, hx4), 1)) + hx4dup = _upsample_like(hx4d, hx3) + + hx3d = self.stage3d(torch.cat((hx4dup, hx3), 1)) + hx3dup = _upsample_like(hx3d, hx2) + + hx2d = self.stage2d(torch.cat((hx3dup, hx2), 1)) + hx2dup = _upsample_like(hx2d, hx1) + + hx1d = self.stage1d(torch.cat((hx2dup, hx1), 1)) + + # side output + d1 = self.side1(hx1d) + d1 = _upsample_like(d1, x) + + d2 = self.side2(hx2d) + d2 = _upsample_like(d2, x) + + d3 = self.side3(hx3d) + d3 = _upsample_like(d3, x) + + d4 = self.side4(hx4d) + d4 = _upsample_like(d4, x) + + d5 = self.side5(hx5d) + d5 = _upsample_like(d5, x) + + d6 = self.side6(hx6) + d6 = _upsample_like(d6, x) + + return [ + F.sigmoid(d1), + F.sigmoid(d2), + F.sigmoid(d3), + F.sigmoid(d4), + F.sigmoid(d5), + F.sigmoid(d6), + ], [hx1d, hx2d, hx3d, hx4d, hx5d, hx6] diff --git a/faceless/processors/face_swapper.py b/faceless/processors/face_swapper.py index 5aeee81..f800529 100644 --- a/faceless/processors/face_swapper.py +++ b/faceless/processors/face_swapper.py @@ -155,7 +155,7 @@ def process_images(source_image, target_images, output_frames_path): raise Exception("process frame failed") write_image(output_filepath, output_vision_frame) -def process_frames(source_image, target_frames_path: str, queue_payloads: List[str], output_frames_path: str): +def process_frames(source_image, target_frames_dir: str, queue_payloads: List[str]): source_frame = tensor_to_vision_frame(source_image) if source_frame is None: raise Exception("cannot read source image") @@ -166,8 +166,7 @@ def process_frames(source_image, target_frames_path: str, queue_payloads: List[s 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_path, frame_filename) - output_filepath = os.path.join(output_frames_path, frame_filename) + frame_filepath = os.path.join(target_frames_dir, frame_filename) target_vision_frame = read_image(frame_filepath) if target_vision_frame is None: @@ -175,17 +174,17 @@ def process_frames(source_image, target_frames_path: str, queue_payloads: List[s output_vision_frame = process_frame(source_face, source_frame, target_vision_frame) if output_vision_frame is None: raise Exception("process frame failed") - write_image(output_filepath, output_vision_frame) + write_image(frame_filepath, output_vision_frame) -def swap_video(source_image, target_frames_path: str, output_frames_path: str): - frames_filenames = os.listdir(target_frames_path) +def swap_video(source_image, target_frames_dir: str): + frames_filenames = os.listdir(target_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, source_image, target_frames_path, pick_queue(queue, queue_per_future), output_frames_path) + future = executor.submit(process_frames, source_image, target_frames_dir, pick_queue(queue, queue_per_future)) futures.append(future) for future_done in as_completed(futures): future_done.result() diff --git a/faceless/typing.py b/faceless/typing.py index 9c2518f..fc2d582 100644 --- a/faceless/typing.py +++ b/faceless/typing.py @@ -38,7 +38,11 @@ Translation = numpy.ndarray[Any, Any] FrameFormat = Literal['jpg', 'png', 'bmp'] FacelessVideo = TypedDict('FacelessVideo', { + # raw vidoe file path 'video_path': str, + # frames dir + 'frames_dir': str, + # output vidoe file path 'output_path': str, 'resolution': Resolution, 'fps': Fps,