diff --git a/birefnet/models/backbones/swin_v1.py b/birefnet/models/backbones/swin_v1.py index 0d5489c..3a360ff 100644 --- a/birefnet/models/backbones/swin_v1.py +++ b/birefnet/models/backbones/swin_v1.py @@ -381,8 +381,8 @@ class BasicLayer(nn.Module): # calculate attention mask for SW-MSA # Turn int to torch.tensor for the compatiability with torch.compile in PyTorch 2.5. - Hp = torch.ceil(torch.tensor(H) / self.window_size).int() * self.window_size - Wp = torch.ceil(torch.tensor(W) / self.window_size).int() * self.window_size + Hp = torch.ceil(torch.tensor(H) / self.window_size).to(torch.int64) * self.window_size + Wp = torch.ceil(torch.tensor(W) / self.window_size).to(torch.int64) * self.window_size img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1 h_slices = (slice(0, -self.window_size), slice(-self.window_size, -self.shift_size), diff --git a/birefnetNode.py b/birefnetNode.py index fb88da5..4bce6b2 100644 --- a/birefnetNode.py +++ b/birefnetNode.py @@ -9,8 +9,7 @@ import folder_paths from birefnet.models.birefnet import BiRefNet from birefnet_old.models.birefnet import BiRefNet as OldBiRefNet from birefnet.utils import check_state_dict -from .util import tensor_to_pil, apply_mask_to_image, normalize_mask - +from .util import tensor_to_pil, apply_mask_to_image, normalize_mask, refine_foreground deviceType = model_management.get_torch_device().type models_dir_key = "birefnet" @@ -58,21 +57,26 @@ def download_birefnet_model(model_name): ) download_models(model_root, model_urls) +class ImagePreprocessor(): + def __init__(self, resolution) -> None: + self.transform_image = transforms.Compose([ + transforms.Resize(resolution), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ]) + self.transform_image_old = transforms.Compose([ + transforms.Resize(resolution), + transforms.ToTensor(), + transforms.Normalize([0.5, 0.5, 0.5], [1.0, 1.0, 1.0]), + ]) -proc_img = transforms.Compose( - [ - transforms.Resize((1024, 1024)), - transforms.ToTensor(), - transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), - ] -) -old_proc_img = transforms.Compose( - [ - transforms.Resize((1024, 1024)), - transforms.ToTensor(), - transforms.Normalize([0.5, 0.5, 0.5], [1.0, 1.0, 1.0]), - ] - ) + def proc(self, image) -> torch.Tensor: + image = self.transform_image(image) + return image + + def old_proc(self, image) -> torch.Tensor: + image = self.transform_image_old(image) + return image VERSION = ["old", "v1"] old_models_name = ["BiRefNet-DIS_ep580.pth", "BiRefNet-ep480.pth"] @@ -160,7 +164,101 @@ class LoadRembgByBiRefNetModel: return [(biRefNet_model, version)] -class RembgByBiRefNet: +class RembgByBiRefNetAdvanced: + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("BiRefNetMODEL",), + "images": ("IMAGE",), + "width": ("INT", + { + "default": 1024, + "min": 0, + "max": 16384, + "tooltip": "The width of the preprocessed image, does not affect the final output image size" + }), + "height": ("INT", + { + "default": 1024, + "min": 0, + "max": 16384, + "tooltip": "The height of the preprocessed image, does not affect the final output image size" + }), + "upscale_method": (["bislerp", "nearest-exact", "bilinear", "area", "bicubic"], + { + "default": "bilinear", + "tooltip": "Interpolation method for post-processing mask" + }), + "blur_size": ("INT", {"default": 91, "min": 1, "max": 255, "step": 2, }), + "blur_size_two": ("INT", {"default": 7, "min": 1, "max": 255, "step": 2, }), + "fill_color": ("BOOLEAN", {"default": False}), + "color": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFF, "step": 1, "display": "color"}), + } + } + + RETURN_TYPES = ("IMAGE", "MASK",) + RETURN_NAMES = ("image", "mask",) + FUNCTION = "rem_bg" + CATEGORY = "rembg/BiRefNet" + + def rem_bg(self, model, images, upscale_method='bilinear', width=1024, height=1024, blur_size=91, blur_size_two=7, fill_color=False, color=None): + model, version = model + model_device_type = next(model.parameters()).device.type + _images = [] + _masks = [] + + for image in images: + h, w, c = image.shape + pil_image = tensor_to_pil(image) + + image_preproc = ImagePreprocessor(resolution=(width, height)) + if VERSION[0] == version: + im_tensor = image_preproc.old_proc(pil_image).unsqueeze(0) + else: + im_tensor = image_preproc.proc(pil_image).unsqueeze(0) + + del image_preproc + + with torch.no_grad(): + mask = model(im_tensor.to(model_device_type))[-1].sigmoid().cpu() + + # 遮罩大小需还原为与原图一致 + mask = comfy.utils.common_upscale(mask, w, h, upscale_method, "disabled") + + # (1, 1, h, w) + mask = normalize_mask(mask) + # (c, h, w) => (c, h, w) + _image_masked = refine_foreground(image.permute(2, 0, 1), mask.squeeze(0), r1=blur_size, r2=blur_size_two).squeeze(0) + # (c, h, w) => (h, w, c) + _image_masked = _image_masked.permute(1, 2, 0) + if fill_color and color is not None: + r = torch.full([h, w, 1], ((color >> 16) & 0xFF) / 0xFF) + g = torch.full([h, w, 1], ((color >> 8) & 0xFF) / 0xFF) + b = torch.full([h, w, 1], (color & 0xFF) / 0xFF) + # (h, w, 3) + background_color = torch.cat((r, g, b), dim=-1) + # (h, w, 1) + apply_mask = mask.squeeze(0).permute(1, 2, 0).expand_as(_image_masked) + _image_masked = _image_masked * apply_mask + background_color * (1 - apply_mask) + # (h, w, 3)=>(1, h, w,3) + image = _image_masked.unsqueeze(0) + del background_color, apply_mask + else: + # image的非mask对应部分设为透明 => (1, h, w, 4) + image = apply_mask_to_image(_image_masked.cpu(), mask.cpu()) + + _images.append(image) + _masks.append(mask.squeeze(0)) + + out_images = torch.cat(_images, dim=0) + out_masks = torch.cat(_masks, dim=0) + + return out_images, out_masks + + +class RembgByBiRefNet(RembgByBiRefNetAdvanced): @classmethod def INPUT_TYPES(cls): @@ -177,47 +275,19 @@ class RembgByBiRefNet: CATEGORY = "rembg/BiRefNet" def rem_bg(self, model, images): - model, version = model - model_device_type = next(model.parameters()).device.type - _images = [] - _masks = [] - - for image in images: - h, w, c = image.shape - pil_image = tensor_to_pil(image) - - if VERSION[0] == version: - im_tensor = old_proc_img(pil_image).unsqueeze(0) - else: - im_tensor = proc_img(pil_image).unsqueeze(0) - - with torch.no_grad(): - mask = model(im_tensor.to(model_device_type))[-1].sigmoid().cpu() - - # 遮罩大小需还原为与原图一致 - mask = comfy.utils.common_upscale(mask, w, h, 'bilinear', "disabled") - - mask = normalize_mask(mask) - # image的非mask对应部分设为透明 - image = apply_mask_to_image(image.cpu(), mask.cpu()) - - _images.append(image) - _masks.append(mask.squeeze(0)) - - out_images = torch.cat(_images, dim=0) - out_masks = torch.cat(_masks, dim=0) - - return out_images, out_masks + return super().rem_bg(model, images) NODE_CLASS_MAPPINGS = { "AutoDownloadBiRefNetModel": AutoDownloadBiRefNetModel, "LoadRembgByBiRefNetModel": LoadRembgByBiRefNetModel, "RembgByBiRefNet": RembgByBiRefNet, + "RembgByBiRefNetAdvanced": RembgByBiRefNetAdvanced, } NODE_DISPLAY_NAME_MAPPINGS = { "AutoDownloadBiRefNetModel": "AutoDownloadBiRefNetModel", "LoadRembgByBiRefNetModel": "LoadRembgByBiRefNetModel", "RembgByBiRefNet": "RembgByBiRefNet", + "RembgByBiRefNetAdvanced": "RembgByBiRefNetAdvanced", } diff --git a/birefnet_old/models/backbones/swin_v1.py b/birefnet_old/models/backbones/swin_v1.py index 4ba6a1e..62a7aea 100644 --- a/birefnet_old/models/backbones/swin_v1.py +++ b/birefnet_old/models/backbones/swin_v1.py @@ -381,8 +381,8 @@ class BasicLayer(nn.Module): # calculate attention mask for SW-MSA # Turn int to torch.tensor for the compatiability with torch.compile in PyTorch 2.5. - Hp = torch.ceil(torch.tensor(H) / self.window_size).int() * self.window_size - Wp = torch.ceil(torch.tensor(W) / self.window_size).int() * self.window_size + Hp = torch.ceil(torch.tensor(H) / self.window_size).to(torch.int64) * self.window_size + Wp = torch.ceil(torch.tensor(W) / self.window_size).to(torch.int64) * self.window_size img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1 h_slices = (slice(0, -self.window_size), slice(-self.window_size, -self.shift_size), diff --git a/pyproject.toml b/pyproject.toml index 194b4cc..0baa12a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_birefnet_ll" description = "Sync with version of BiRefNet. NODES:AutoDownloadBiRefNetModel, LoadRembgByBiRefNetModel, RembgByBiRefNet." -version = "1.0.5" +version = "1.0.6" license = {file = "LICENSE"} dependencies = ["numpy", "opencv-python", "timm"] diff --git a/util.py b/util.py index a515eda..2fcb654 100644 --- a/util.py +++ b/util.py @@ -1,6 +1,7 @@ import numpy as np import torch from PIL import Image +import torchvision.transforms.v2 as T def tensor_to_pil(image): @@ -11,6 +12,41 @@ def pil_to_tensor(image): return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) +def refine_foreground(image_tensor, mask_tensor, r1=90, r2=7): + if r1 % 2 == 0: + r1 += 1 + + if r2 % 2 == 0: + r2 += 1 + + estimated_foreground = FB_blur_fusion_foreground_estimator_2(image_tensor, mask_tensor, r1=r1, r2=r2) + return estimated_foreground + + +def FB_blur_fusion_foreground_estimator_2(image_tensor, alpha_tensor, r1=90, r2=7): + # https://github.com/Photoroom/fast-foreground-estimation + if alpha_tensor.dim() == 3: + alpha_tensor = alpha_tensor.unsqueeze(0) # Add batch + F, blur_B = FB_blur_fusion_foreground_estimator(image_tensor, image_tensor, image_tensor, alpha_tensor, r=r1) + return FB_blur_fusion_foreground_estimator(image_tensor, F, blur_B, alpha_tensor, r=r2)[0] + + +def FB_blur_fusion_foreground_estimator(image_tensor, F_tensor, B_tensor, alpha_tensor, r=90): + if image_tensor.dim() == 3: + image_tensor = image_tensor.unsqueeze(0) + + blurred_alpha = T.functional.gaussian_blur(alpha_tensor, r) + + blurred_FA = T.functional.gaussian_blur(F_tensor * alpha_tensor, r) + blurred_F = blurred_FA / (blurred_alpha + 1e-5) + + blurred_B1A = T.functional.gaussian_blur(B_tensor * (1 - alpha_tensor), r) + blurred_B = blurred_B1A / ((1 - blurred_alpha) + 1e-5) + F_tensor = blurred_F + alpha_tensor * (image_tensor - alpha_tensor * blurred_F - (1 - alpha_tensor) * blurred_B) + F_tensor = torch.clamp(F_tensor, 0, 1) + return F_tensor, blurred_B + + def apply_mask_to_image(image, mask): """ Apply a mask to an image and set non-masked parts to transparent.