feat: support remove background

This commit is contained in:
Jeffrey Wu
2024-05-13 17:44:14 +08:00
parent 9ac5fdc643
commit 3bb861b608
13 changed files with 668 additions and 49 deletions
+3 -3
View File
@@ -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' ]:
+10
View File
@@ -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)
+7
View File
@@ -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)",
}
+1 -1
View File
@@ -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)):
+7 -6
View File
@@ -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,)
+13 -12
View File
@@ -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,)
+104
View File
@@ -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
+10 -8
View File
@@ -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 ()
+2 -12
View File
@@ -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,)
@@ -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,)
+457
View File
@@ -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]
+6 -7
View File
@@ -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()
+4
View File
@@ -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,