use cv2 to implement BlurFusionForegroundEstimation. https://github.com/lldacing/ComfyUI_BiRefNet_ll/issues/23
This commit is contained in:
+8
-5
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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, }),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user