From 8a71b9e345607ce18d5366186cec04dc335fa65a Mon Sep 17 00:00:00 2001 From: Hacker 17082006 Date: Wed, 7 Feb 2024 19:55:33 +0700 Subject: [PATCH] Add BRIAI rembg v1.4 --- README.md | 2 +- __init__.py | 108 +---- comfyui_vidmatt/__init__.py | 0 comfyui_vidmatt/briaai_rembg/__init__.py | 68 +++ comfyui_vidmatt/briaai_rembg/arch.py | 454 ++++++++++++++++++ .../robust_video_matting/__init__.py | 56 +++ comfyui_vidmatt/utils.py | 60 +++ 7 files changed, 650 insertions(+), 98 deletions(-) create mode 100644 comfyui_vidmatt/__init__.py create mode 100644 comfyui_vidmatt/briaai_rembg/__init__.py create mode 100644 comfyui_vidmatt/briaai_rembg/arch.py create mode 100644 comfyui_vidmatt/robust_video_matting/__init__.py create mode 100644 comfyui_vidmatt/utils.py diff --git a/README.md b/README.md index b812a7c..f6a8b8b 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -A minimalistic implementation of [Robust Video Matting (RVM)](https://github.com/PeterL1n/RobustVideoMatting/) in ComfyUI +A minimalistic implementation of [Robust Video Matting (RVM)](https://github.com/PeterL1n/RobustVideoMatting/) and in ComfyUI [Example workflow](./example_matting_workflow.json) diff --git a/__init__.py b/__init__.py index b1854b2..f227cac 100644 --- a/__init__.py +++ b/__init__.py @@ -1,102 +1,16 @@ -import torch -from einops import rearrange, repeat -import os, yaml -from torch.hub import download_url_to_file, get_dir -from urllib.parse import urlparse -from comfy.model_management import soft_empty_cache, get_torch_device -from PIL import ImageColor +from pathlib import Path +import os, sys -config_path = os.path.join(os.path.dirname(__file__), "./config.yaml") -if os.path.exists(config_path): - config = yaml.load(open(config_path, "r"), Loader=yaml.FullLoader) -else: - raise Exception("config.yaml file is neccessary, plz recreate the config file by downloading it from https://github.com/Fannovel16/ComfyUI-Video-Matting") -CKPTS_PATH = os.path.join(os.path.dirname(__file__), config["ckpts_path"]) +ext_path = Path(__file__).parent +sys.path.insert(0, str(ext_path.resolve())) +for model in os.listdir((ext_path / "comfyui_vidmatt").resolve()): + model_path = (ext_path / "comfyui_vidmatt" / model).resolve() + sys.path.insert(0, str(model_path)) -def auto_downsample_ratio(h, w): - """ - Automatically find a downsample ratio so that the largest side of the resolution be 512px. - """ - return min(512 / max(h, w), 1) -def load_file_from_url(url, model_dir=None, progress=True, file_name=None): - """Load file form http url, will download models if necessary. - - Ref:https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py - - Args: - url (str): URL to be downloaded. - model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir. - Default: None. - progress (bool): Whether to show the download progress. Default: True. - file_name (str): The downloaded file name. If None, use the file name in the url. Default: None. - - Returns: - str: The path to the downloaded file. - """ - if model_dir is None: # use the pytorch hub_dir - hub_dir = get_dir() - model_dir = os.path.join(hub_dir, 'checkpoints') - - os.makedirs(model_dir, exist_ok=True) - - parts = urlparse(url) - file_name = os.path.basename(parts.path) - if file_name is not None: - file_name = file_name - cached_file = os.path.abspath(os.path.join(model_dir, file_name)) - if not os.path.exists(cached_file): - print(f'Downloading: "{url}" to {cached_file}\n') - download_url_to_file(url, cached_file, hash_prefix=None, progress=progress) - return cached_file - -download_url_template = "https://github.com/PeterL1n/RobustVideoMatting/releases/download/v1.0.0/rvm_{backbone}_{dtype}.torchscript" -device = get_torch_device() -class RobustVideoMatting: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "video_frames": ("IMAGE",), - "backbone": (["mobilenetv3", "resnet50"], {"default": "resnet50"}), - "fp16": ("BOOLEAN", {"default": True}), - "bg_color": ("STRING", {"default": "green"}), - "batch_size": ("INT", {"min": 1, "max": 64, "default": 4}) - } - } - - RETURN_TYPES = ("IMAGE", "MASK") - FUNCTION = "matting" - CATEGORY = "Video Matting/Robust Video Matting" - - def matting(self, video_frames, backbone, fp16, bg_color, batch_size): - model_path = load_file_from_url(download_url_template.format(backbone=backbone, dtype="fp16" if fp16 else "fp32"), model_dir=CKPTS_PATH) - model = torch.jit.load(model_path, map_location="cpu") - model.to(device) - video_frames = rearrange(video_frames, "n h w c -> n c h w") - bg_color = torch.Tensor(ImageColor.getrgb(bg_color)[:3]).to(device).float() / 255. - bg_color = repeat(bg_color, "c -> n c 1 1", n=batch_size) - if fp16: - model.half() - bg_color.half() - model = torch.jit.freeze(model) - orig_num_frames = video_frames.shape[0] - pad_frames = repeat(video_frames[-1:], "1 c h w -> n c h w", n=batch_size - (orig_num_frames % batch_size)) - video_frames = torch.cat([video_frames, pad_frames], dim=0) - rec, fgrs, masks = [None] * 4, [], [] - for i in range(video_frames.shape[0] // batch_size): - input = video_frames[i*batch_size:(i+1)*batch_size].to(device) - if fp16: - input = input.half() - fgr, pha, *rec = model(input, *rec, auto_downsample_ratio(*video_frames.shape[2:])) - mask = pha.gt(0) - fgr = fgr * mask + bg_color * ~mask - fgrs.append(fgr.float()) - masks.append(mask.float()) - fgrs = rearrange(torch.cat(fgrs, dim=0), "n c h w -> n h w c")[:orig_num_frames].cpu().detach() - masks = torch.cat(masks, dim=0)[:orig_num_frames].cpu().detach() - soft_empty_cache() - return (fgrs, masks) +from comfyui_vidmatt.robust_video_matting import RobustVideoMatting +from comfyui_vidmatt.briaai_rembg import BriaaiRembg NODE_CLASS_MAPPINGS = { - "Robust Video Matting": RobustVideoMatting + "Robust Video Matting": RobustVideoMatting, + "BRIAAI Matting": BriaaiRembg } \ No newline at end of file diff --git a/comfyui_vidmatt/__init__.py b/comfyui_vidmatt/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/comfyui_vidmatt/briaai_rembg/__init__.py b/comfyui_vidmatt/briaai_rembg/__init__.py new file mode 100644 index 0000000..06c3e90 --- /dev/null +++ b/comfyui_vidmatt/briaai_rembg/__init__.py @@ -0,0 +1,68 @@ +import torch, os +from PIL import Image +from comfyui_vidmatt.briaai_rembg.arch import BriaRMBG +import torch.nn.functional as F +from torchvision.transforms.functional import normalize + +import torch +from einops import rearrange, repeat +from comfy.model_management import soft_empty_cache, get_torch_device +from PIL import ImageColor + +from comfyui_vidmatt.utils import CKPTS_PATH, load_file_from_url, prepare_frames_color + +def auto_downsample_ratio(h, w): + """ + Automatically find a downsample ratio so that the largest side of the resolution be 512px. + """ + return min(512 / max(h, w), 1) + +download_url = "https://huggingface.co/briaai/RMBG-1.4/resolve/main/model.pth" +device = get_torch_device() +model_input_size = [1024,1024] + +class BriaaiRembg: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "video_frames": ("IMAGE",), + "version": (["v1.4"], {"default": "v1.4"}), + "fp16": ("BOOLEAN", {"default": True}), + "bg_color": ("STRING", {"default": "green"}), + "batch_size": ("INT", {"min": 1, "max": 64, "default": 4}) + } + } + + RETURN_TYPES = ("IMAGE", "MASK") + FUNCTION = "matting" + CATEGORY = "Video Matting" + + + def matting(self, video_frames, version, fp16, bg_color, batch_size): + model_path = load_file_from_url(download_url, file_name=f"briaai_rmbg_{version}.pth", model_dir=CKPTS_PATH) + model = BriaRMBG() + model.load_state_dict(torch.load(model_path, map_location=device)) + model.to(device).eval() + + video_frames, orig_num_frames, bg_color = prepare_frames_color(video_frames, bg_color, batch_size) + bg_color.to(device) + video_frames = F.interpolate(video_frames, size=model_input_size, mode='bilinear') + if fp16: + model.half() + bg_color.half() + + fgrs, masks = [], [] + for i in range(video_frames.shape[0] // batch_size): + input = video_frames[i*batch_size:(i+1)*batch_size].to(device) + if fp16: + input = input.half() + mask = model(normalize(input,[0.5,0.5,0.5],[1.0,1.0,1.0]))[0][0] + mask = (mask-mask.min())/(mask.max()-mask.min()) #This is sharp enough + fgr = input * mask + bg_color * (1 - mask) + fgrs.append(fgr.float()) + masks.append(mask.float()) + fgrs = rearrange(torch.cat(fgrs, dim=0), "n c h w -> n h w c")[:orig_num_frames].cpu().detach() + masks = torch.cat(masks, dim=0)[:orig_num_frames].cpu().detach() + soft_empty_cache() + return (fgrs, masks) \ No newline at end of file diff --git a/comfyui_vidmatt/briaai_rembg/arch.py b/comfyui_vidmatt/briaai_rembg/arch.py new file mode 100644 index 0000000..9187700 --- /dev/null +++ b/comfyui_vidmatt/briaai_rembg/arch.py @@ -0,0 +1,454 @@ +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 + +## upsample tensor 'src' to have the same spatial size with tensor 'tar' +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,in_ch=3,out_ch=1): + super(BriaRMBG,self).__init__() + + 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/comfyui_vidmatt/robust_video_matting/__init__.py b/comfyui_vidmatt/robust_video_matting/__init__.py new file mode 100644 index 0000000..efc773e --- /dev/null +++ b/comfyui_vidmatt/robust_video_matting/__init__.py @@ -0,0 +1,56 @@ +import torch +from einops import rearrange, repeat +from comfy.model_management import soft_empty_cache, get_torch_device +from PIL import ImageColor + +from comfyui_vidmatt.utils import CKPTS_PATH, load_file_from_url, prepare_frames_color + +def auto_downsample_ratio(h, w): + """ + Automatically find a downsample ratio so that the largest side of the resolution be 512px. + """ + return min(512 / max(h, w), 1) + +download_url_template = "https://github.com/PeterL1n/RobustVideoMatting/releases/download/v1.0.0/rvm_{backbone}_{dtype}.torchscript" +device = get_torch_device() +class RobustVideoMatting: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "video_frames": ("IMAGE",), + "backbone": (["mobilenetv3", "resnet50"], {"default": "resnet50"}), + "fp16": ("BOOLEAN", {"default": True}), + "bg_color": ("STRING", {"default": "green"}), + "batch_size": ("INT", {"min": 1, "max": 64, "default": 4}) + } + } + + RETURN_TYPES = ("IMAGE", "MASK") + FUNCTION = "matting" + CATEGORY = "Video Matting" + + def matting(self, video_frames, backbone, fp16, bg_color, batch_size): + model_path = load_file_from_url(download_url_template.format(backbone=backbone, dtype="fp16" if fp16 else "fp32"), model_dir=CKPTS_PATH) + model = torch.jit.load(model_path, map_location="cpu") + model.to(device) + video_frames, orig_num_frames, bg_color = prepare_frames_color(video_frames, bg_color, batch_size) + bg_color.to(device) + if fp16: + model.half() + bg_color.half() + model = torch.jit.freeze(model) + rec, fgrs, masks = [None] * 4, [], [] + for i in range(video_frames.shape[0] // batch_size): + input = video_frames[i*batch_size:(i+1)*batch_size].to(device) + if fp16: + input = input.half() + fgr, pha, *rec = model(input, *rec, auto_downsample_ratio(*video_frames.shape[2:])) + mask = pha.gt(0) #Remove blur + fgr = fgr * mask + bg_color * ~mask + fgrs.append(fgr.float()) + masks.append(mask.float()) + fgrs = rearrange(torch.cat(fgrs, dim=0), "n c h w -> n h w c")[:orig_num_frames].cpu().detach() + masks = torch.cat(masks, dim=0)[:orig_num_frames].cpu().detach() + soft_empty_cache() + return (fgrs, masks) \ No newline at end of file diff --git a/comfyui_vidmatt/utils.py b/comfyui_vidmatt/utils.py new file mode 100644 index 0000000..5801a5e --- /dev/null +++ b/comfyui_vidmatt/utils.py @@ -0,0 +1,60 @@ +import os, yaml +from torch.hub import download_url_to_file, get_dir +from urllib.parse import urlparse +from einops import rearrange, repeat +import torch +from PIL import ImageColor + +config_path = os.path.join(os.path.dirname(__file__), "../config.yaml") +if os.path.exists(config_path): + config = yaml.load(open(config_path, "r"), Loader=yaml.FullLoader) +else: + raise Exception("config.yaml file is neccessary, plz recreate the config file by downloading it from https://github.com/Fannovel16/ComfyUI-Video-Matting") +CKPTS_PATH = os.path.join(os.path.join(os.path.dirname(__file__), '..'), config["ckpts_path"]) + +def auto_downsample_ratio(h, w): + """ + Automatically find a downsample ratio so that the largest side of the resolution be 512px. + """ + return min(512 / max(h, w), 1) + +def load_file_from_url(url, model_dir=None, progress=True, file_name=None): + """Load file form http url, will download models if necessary. + + Ref:https://github.com/1adrianb/face-alignment/blob/master/face_alignment/utils.py + + Args: + url (str): URL to be downloaded. + model_dir (str): The path to save the downloaded model. Should be a full path. If None, use pytorch hub_dir. + Default: None. + progress (bool): Whether to show the download progress. Default: True. + file_name (str): The downloaded file name. If None, use the file name in the url. Default: None. + + Returns: + str: The path to the downloaded file. + """ + if model_dir is None: # use the pytorch hub_dir + hub_dir = get_dir() + model_dir = os.path.join(hub_dir, 'checkpoints') + + os.makedirs(model_dir, exist_ok=True) + + parts = urlparse(url) + if file_name is None: + file_name = os.path.basename(parts.path) + cached_file = os.path.abspath(os.path.join(model_dir, file_name)) + if not os.path.exists(cached_file): + print(f'Downloading: "{url}" to {cached_file}\n') + download_url_to_file(url, cached_file, hash_prefix=None, progress=progress) + return cached_file + +def prepare_frames_color(video_frames, bg_color, batch_size): + orig_num_frames = video_frames.shape[0] + video_frames = rearrange(video_frames, "n h w c -> n c h w") + pad_frames = repeat(video_frames[-1:], "1 c h w -> n c h w", n=batch_size - (orig_num_frames % batch_size)) + video_frames = torch.cat([video_frames, pad_frames], dim=0) + + bg_color = torch.Tensor(ImageColor.getrgb(bg_color)[:3]).float() / 255. + bg_color = repeat(bg_color, "c -> n c 1 1", n=batch_size) + + return video_frames, orig_num_frames, bg_color