use cv2 to implement BlurFusionForegroundEstimation. https://github.com/lldacing/ComfyUI_BiRefNet_ll/issues/23

This commit is contained in:
刘雪峰
2025-05-03 12:56:55 +08:00
parent 617827cee9
commit c8765cfaf9
6 changed files with 129 additions and 34 deletions
+8 -5
View File
@@ -32,14 +32,18 @@ class Config:
'General-2K': 'DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4+DIS-TR+TR-HRSOD+TE-HRSOD+TR-HRS10K+TE-HRS10K+TR-UHRSD+TE-UHRSD+TR-P3M-10k+TE-P3M-500-P+TR-humans+DIS-VD-ori', # datasets_all
'Matting': 'TR-P3M-10k+TE-P3M-500-NP+TR-humans+TR-Distrinctions-646', # datasets_all
}[self.task]
self.prompt4loc = ['dense', 'sparse'][0]
# Data settings
self.size = (1024, 1024) if self.task not in ['General-2K'] else (2560, 1440) # wid, hei. Can be overwritten by dynamic_size in training.
self.dynamic_size = [None, ((512-256, 2048+256), (512-256, 2048+256))][0] # wid, hei. It might cause errors in using compile.
self.background_color_synthesis = False # whether to use pure bg color to replace the original backgrounds.
# Faster-Training settings
self.load_all = False # Turn it on/off by your case. It may consume a lot of CPU memory. And for multi-GPU (N), it would cost N times the CPU memory to load the data.
self.load_all = False and self.dynamic_size is None # Turn it on/off by your case. It may consume a lot of CPU memory. And for multi-GPU (N), it would cost N times the CPU memory to load the data.
self.compile = True # 1. Trigger CPU memory leak in some extend, which is an inherent problem of PyTorch.
# Machines with > 70GB CPU memory can run the whole training on DIS5K with default setting.
# 2. Higher PyTorch version may fix it: https://github.com/pytorch/pytorch/issues/119607.
# 3. But compile in Pytorch > 2.0.1 seems to bring no acceleration for training.
# 3. But compile in 2.0.1 < Pytorch < 2.5.0 seems to bring no acceleration for training.
self.precisionHigh = True
# MODEL settings
@@ -67,7 +71,6 @@ class Config:
}[self.task]
][1] # choose 0 to skip
self.lr = (1e-4 if 'DIS5K' in self.task else 1e-5) * math.sqrt(self.batch_size / 4) # DIS needs high lr to converge faster. Adapt the lr linearly
self.size = (1024, 1024) if self.task not in ['General-2K'] else (2560, 1440) # wid, hei
self.num_workers = max(4, self.batch_size) # will be decrease to min(it, batch_size) at the initialization of the data_loader
# Backbone settings
@@ -105,7 +108,7 @@ class Config:
][0]
# TRAINING settings - inactive
self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4]
self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4 if not self.background_color_synthesis else 1]
self.optimizer = ['Adam', 'AdamW'][1]
self.lr_decay_epochs = [1e5] # Set to negative N to decay the lr in the last N-th epoch.
self.lr_decay_rate = 0.5
+67 -15
View File
@@ -1,4 +1,6 @@
import os
import random
import numpy as np
import cv2
from tqdm import tqdm
from PIL import Image
@@ -32,12 +34,10 @@ class_labels_TR_sorted = _class_labels_TR_sorted.split(', ')
class MyData(data.Dataset):
def __init__(self, datasets, image_size, is_train=True):
self.size_train = image_size
self.size_test = image_size
self.keep_size = not config.size
self.data_size = config.size
def __init__(self, datasets, data_size, is_train=True):
# data_size is None when using dynamic_size or data_size is manually set to None (for inference in the original size).
self.is_train = is_train
self.data_size = data_size
self.load_all = config.load_all
self.device = config.device
valid_extensions = ['.png', '.jpg', '.PNG', '.JPG', '.JPEG']
@@ -45,14 +45,12 @@ class MyData(data.Dataset):
if self.is_train and config.auxiliary_classification:
self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
self.transform_image = transforms.Compose([
transforms.Resize(self.data_size[::-1]),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
][self.load_all or self.keep_size:])
])
self.transform_label = transforms.Compose([
transforms.Resize(self.data_size[::-1]),
transforms.ToTensor(),
][self.load_all or self.keep_size:])
])
dataset_root = os.path.join(config.data_root_dir, config.task)
# datasets can be a list of different datasets for training on combined sets.
self.image_paths = []
@@ -83,8 +81,8 @@ class MyData(data.Dataset):
self.class_labels_loaded = []
# for image_path, label_path in zip(self.image_paths, self.label_paths):
for image_path, label_path in tqdm(zip(self.image_paths, self.label_paths), total=len(self.image_paths)):
_image = path_to_image(image_path, size=config.size, color_type='rgb')
_label = path_to_image(label_path, size=config.size, color_type='gray')
_image = path_to_image(image_path, size=self.data_size, color_type='rgb')
_label = path_to_image(label_path, size=self.data_size, color_type='gray')
self.images_loaded.append(_image)
self.labels_loaded.append(_label)
self.class_labels_loaded.append(
@@ -92,25 +90,57 @@ class MyData(data.Dataset):
)
def __getitem__(self, index):
if self.load_all:
image = self.images_loaded[index]
label = self.labels_loaded[index]
class_label = self.class_labels_loaded[index] if self.is_train and config.auxiliary_classification else -1
else:
image = path_to_image(self.image_paths[index], size=config.size, color_type='rgb')
label = path_to_image(self.label_paths[index], size=config.size, color_type='gray')
image = path_to_image(self.image_paths[index], size=self.data_size, color_type='rgb')
label = path_to_image(self.label_paths[index], size=self.data_size, color_type='gray')
class_label = self.cls_name2id[self.label_paths[index].split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
# loading image and label
if self.is_train:
if config.background_color_synthesis:
image.putalpha(label)
array_image = np.array(image)
array_foreground = array_image[:, :, :3].astype(np.float32)
array_mask = (array_image[:, :, 3:] / 255).astype(np.float32)
array_background = np.zeros_like(array_foreground)
choice = random.random()
if choice < 0.4:
# Black/Gray/White backgrounds
array_background[:, :, :] = random.randint(0, 255)
elif choice < 0.8:
# Background color that similar to the foreground object. Hard negative samples.
foreground_pixel_number = np.sum(array_mask > 0)
color_foreground_mean = np.mean(array_foreground * array_mask, axis=(0, 1)) * (np.prod(array_foreground.shape[:2]) / foreground_pixel_number)
color_up_or_down = random.choice((-1, 1))
# Up or down for 20% range from 255 or 0, respectively.
color_foreground_mean += (255 - color_foreground_mean if color_up_or_down == 1 else color_foreground_mean) * (random.random() * 0.2) * color_up_or_down
array_background[:, :, :] = color_foreground_mean
else:
# Any color
for idx_channel in range(3):
array_background[:, :, idx_channel] = random.randint(0, 255)
array_foreground_background = array_foreground * array_mask + array_background * (1 - array_mask)
image = Image.fromarray(array_foreground_background.astype(np.uint8))
image, label = preproc(image, label, preproc_methods=config.preproc_methods)
# else:
# if _label.shape[0] > 2048 or _label.shape[1] > 2048:
# _image = cv2.resize(_image, (2048, 2048), interpolation=cv2.INTER_LINEAR)
# _label = cv2.resize(_label, (2048, 2048), interpolation=cv2.INTER_LINEAR)
image, label = self.transform_image(image), self.transform_label(label)
# At present, we use fixed sizes in inference, instead of consistent dynamic size with training.
if self.is_train:
if config.dynamic_size is None:
image, label = self.transform_image(image), self.transform_label(label)
else:
size_div_32 = (int(image.size[0] // 32 * 32), int(image.size[1] // 32 * 32))
if image.size != size_div_32:
image = image.resize(size_div_32)
label = label.resize(size_div_32)
image, label = self.transform_image(image), self.transform_label(label)
if self.is_train:
return image, label, class_label
@@ -119,3 +149,25 @@ class MyData(data.Dataset):
def __len__(self):
return len(self.image_paths)
def custom_collate_fn(batch):
if config.dynamic_size:
dynamic_size = tuple(sorted(config.dynamic_size))
dynamic_size_batch = (random.randint(dynamic_size[0][0], dynamic_size[0][1]) // 32 * 32, random.randint(dynamic_size[1][0], dynamic_size[1][1]) // 32 * 32) # select a value randomly in the range of [dynamic_size[0/1][0], dynamic_size[0/1][1]].
data_size = dynamic_size_batch
else:
data_size = config.size
new_batch = []
transform_image = transforms.Compose([
transforms.Resize(data_size[::-1]),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
])
transform_label = transforms.Compose([
transforms.Resize(data_size[::-1]),
transforms.ToTensor(),
])
for image, label, class_label in batch:
new_batch.append((transform_image(image), transform_label(label), class_label))
return data._utils.collate.default_collate(new_batch)
+1 -1
View File
@@ -16,7 +16,7 @@ try:
except Exception:
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
from birefnet.config import Config
from ...config import Config
config = Config()
+5 -4
View File
@@ -26,12 +26,13 @@ def path_to_image(path, size=(1024, 1024), color_type=['rgb', 'gray'][0]):
def check_state_dict(state_dict, unwanted_prefixes=['_orig_mod.', 'module.']):
def check_state_dict(state_dict, unwanted_prefixes=['module.', '_orig_mod.']):
for k, v in list(state_dict.items()):
prefix_length = 0
for unwanted_prefix in unwanted_prefixes:
if k.startswith(unwanted_prefix):
state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
break
if k[prefix_length:].startswith(unwanted_prefix):
prefix_length += len(unwanted_prefix)
state_dict[k[prefix_length:]] = state_dict.pop(k)
return state_dict
+11 -8
View File
@@ -9,7 +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 refine_foreground, filter_mask, add_mask_as_alpha
from .util import filter_mask, add_mask_as_alpha, refine_foreground_pil, tensor_to_pil, pil_to_tensor
deviceType = model_management.get_torch_device().type
models_dir_key = "birefnet"
@@ -264,8 +264,8 @@ class BlurFusionForegroundEstimation:
"required": {
"images": ("IMAGE",),
"masks": ("MASK",),
"blur_size": ("INT", {"default": 91, "min": 1, "max": 255, "step": 2, }),
"blur_size_two": ("INT", {"default": 7, "min": 1, "max": 255, "step": 2, }),
"blur_size": ("INT", {"default": 90, "min": 1, "max": 255, "step": 1, }),
"blur_size_two": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1, }),
"fill_color": ("BOOLEAN", {"default": False}),
"color": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFF, "step": 1, "display": "color"}),
}
@@ -282,16 +282,19 @@ class BlurFusionForegroundEstimation:
if b != masks.shape[0]:
raise ValueError("images and masks must have the same batch size")
image_bchw = images.permute(0, 3, 1, 2)
# image_bchw = images.permute(0, 3, 1, 2)
if masks.dim() == 3:
# (b, h, w) => (b, 1, h, w)
out_masks = masks.unsqueeze(1)
# 需要转成pil用cv2.blur,结果图的背景色比较纯(gaussian_blur的背景色不纯,边缘轮廓线比较重),应用遮罩时不能用点乘,结果可能有边缘轮廓
_image_masked = refine_foreground_pil(tensor_to_pil(images), tensor_to_pil(out_masks.permute(0, 2, 3, 1)))
_image_masked = pil_to_tensor(_image_masked)
# (b, c, h, w)
_image_masked = refine_foreground(image_bchw, out_masks, r1=blur_size, r2=blur_size_two)
# _image_masked = refine_foreground(image_bchw, out_masks, r1=blur_size, r2=blur_size_two)
# (b, c, h, w) => (b, h, w, c)
_image_masked = _image_masked.permute(0, 2, 3, 1)
# _image_masked = _image_masked.permute(0, 2, 3, 1)
if fill_color and color is not None:
r = torch.full([b, h, w, 1], ((color >> 16) & 0xFF) / 0xFF)
g = torch.full([b, h, w, 1], ((color >> 8) & 0xFF) / 0xFF)
@@ -342,8 +345,8 @@ class RembgByBiRefNetAdvanced(GetMaskByBiRefNet, BlurFusionForegroundEstimation)
"default": "bilinear",
"tooltip": "Interpolation method for pre-processing image and 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, }),
"blur_size": ("INT", {"default": 90, "min": 1, "max": 255, "step": 1, }),
"blur_size_two": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1, }),
"fill_color": ("BOOLEAN", {"default": False}),
"color": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFF, "step": 1, "display": "color"}),
"mask_threshold": ("FLOAT", {"default": 0.000, "min": 0.0, "max": 1.0, "step": 0.001, }),
+37 -1
View File
@@ -2,6 +2,7 @@ import numpy as np
import torch
from PIL import Image
import torchvision.transforms.v2 as T
import cv2
def tensor_to_pil(image):
@@ -46,6 +47,40 @@ def FB_blur_fusion_foreground_estimator(image_tensor, F_tensor, B_tensor, alpha_
return F_tensor, blurred_B
### copied and modified image_proc.py
def refine_foreground_pil(image, mask, r1=90, r2=6):
if mask.size != image.size:
mask = mask.resize(image.size)
image = np.array(image) / 255.0
mask = np.array(mask) / 255.0
estimated_foreground = FB_blur_fusion_foreground_estimator_pil_2(image, mask, r1=r1, r2=r2)
image_masked = Image.fromarray((estimated_foreground * 255.0).astype(np.uint8))
return image_masked
def FB_blur_fusion_foreground_estimator_pil_2(image, alpha, r1=90, r2=6):
# Thanks to the source: https://github.com/Photoroom/fast-foreground-estimation
alpha = alpha[:, :, None]
F, blur_B = FB_blur_fusion_foreground_estimator_pil(
image, image, image, alpha, r=r1)
return FB_blur_fusion_foreground_estimator_pil(image, F, blur_B, alpha, r=r2)[0]
def FB_blur_fusion_foreground_estimator_pil(image, F, B, alpha, r=90):
if isinstance(image, Image.Image):
image = np.array(image) / 255.0
blurred_alpha = cv2.blur(alpha, (r, r))[:, :, None]
blurred_FA = cv2.blur(F * alpha, (r, r))
blurred_F = blurred_FA / (blurred_alpha + 1e-5)
blurred_B1A = cv2.blur(B * (1 - alpha), (r, r))
blurred_B = blurred_B1A / ((1 - blurred_alpha) + 1e-5)
F = blurred_F + alpha * (image - alpha * blurred_F - (1 - alpha) * blurred_B)
F = np.clip(F, 0, 1)
return F, blurred_B
def apply_mask_to_image(image, mask):
"""
Apply a mask to an image and set non-masked parts to transparent.
@@ -115,7 +150,8 @@ def add_mask_as_alpha(image, mask):
# 将 mask 扩展为 (b, h, w, 1)
mask = mask[..., None]
image = image * mask
# 不做点乘,可能会有边缘轮廓线
# image = image * mask
# 将 image 和 mask 拼接为 (b, h, w, 4)
image_with_alpha = torch.cat([image, mask], dim=-1)