feat: support remove background
This commit is contained in:
+3
-3
@@ -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' ]:
|
||||
|
||||
@@ -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,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)",
|
||||
}
|
||||
|
||||
@@ -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)):
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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 ()
|
||||
|
||||
@@ -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,)
|
||||
@@ -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]
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user