v1.1: 优化BiRefNet模型加载和推理流程,改进文档结构,移除版本号依赖
This commit is contained in:
@@ -1,37 +1,47 @@
|
||||
# ComfyUI-RemoveBackgroundSuite
|
||||
|
||||
这是一个 ComfyUI 插件,专注于实现各类高质量背景移除功能,支持多种 SOTA 算法和细节处理。
|
||||
> 基于 ComfyUI 的抠图套件,支持多种抠图模型和细节处理方式。
|
||||
|
||||
## 版本信息
|
||||
|
||||
- 当前版本:v1.1
|
||||
- 更新日期:2024-03-21
|
||||
|
||||
## 更新日志
|
||||
|
||||
### v1.1 (2024-03-21)
|
||||
- 优化 BiRefNet 模型加载逻辑,支持 dynamic、HR、HR-matting 模型
|
||||
- 改进 BiRefNetUltraV3_RBS 节点,支持多种细节处理方式
|
||||
- 优化 VITMatte 模型加载和推理流程
|
||||
- 改进文档结构和说明
|
||||
|
||||
### v1.0 (2024-03-20)
|
||||
- 初始版本发布
|
||||
- 支持 BiRefNet 和 TransparentBackground 两种抠图模型
|
||||
- 支持多种细节处理方式
|
||||
|
||||
## 节点说明
|
||||
|
||||
### 1. LoadBiRefNetModel_RBS
|
||||
- **功能**:加载本地 BiRefNet 模型权重。
|
||||
### 1. LoadBiRefNetModelV3_RBS(推荐)
|
||||
- **功能**:自动下载并加载 BiRefNet 最新模型,支持 BiRefNet_dynamic。
|
||||
- **参数**:
|
||||
- `model`:选择本地模型文件(.pth)。
|
||||
- `version`:选择模型版本(BiRefNet-General、RMBG-2.0、BiRefNet_dynamic)。
|
||||
- **输出**:`birefnet_model`(供后续节点使用)
|
||||
|
||||
### 2. LoadBiRefNetModelV2_RBS
|
||||
- **功能**:自动下载并加载 BiRefNet 新版模型(支持 Huggingface 仓库)。
|
||||
- **参数**:
|
||||
- `version`:选择模型版本(如 BiRefNet-General、RMBG-2.0)。
|
||||
- **输出**:`birefnet_model`(供后续节点使用)
|
||||
|
||||
### 3. BiRefNetUltraV2_RBS
|
||||
- **功能**:使用 BiRefNet Ultra V2 进行高质量背景移除。
|
||||
### 2. BiRefNetUltraV3_RBS(推荐)
|
||||
- **功能**:使用 BiRefNet Ultra V3 进行高质量背景移除,支持 dynamic 动态模型。
|
||||
- **参数**:
|
||||
- `image`:输入图片(支持批量)。
|
||||
- `birefnet_model`:已加载的模型。
|
||||
- `detail_method`:细节处理方式(VITMatte、PyMatting、GuidedFilter等)。
|
||||
- `detail_erode`/`detail_dilate`:腐蚀/膨胀参数,影响边缘细节。
|
||||
- `black_point`/`white_point`:黑白场,调整掩码对比度。
|
||||
- `process_detail`:是否进行细节处理。
|
||||
- `device`:推理设备(cuda/cpu)。
|
||||
- `max_megapixels`:最大处理分辨率。
|
||||
- `birefnet_model`:已加载的模型(支持 dynamic)。
|
||||
- `detail_method`、`detail_erode`、`detail_dilate`、`black_point`、`white_point`、`process_detail`、`device`、`max_megapixels`:详见节点界面。
|
||||
- **输出**:
|
||||
- `image`:去背景后的 RGBA 图片
|
||||
- `mask`:前景掩码
|
||||
|
||||
### 4. TransparentBackgroundUltra_RBS
|
||||
### 3. TransparentBackgroundUltra_RBS([WIP] 开发中)
|
||||
|
||||
> ⚠️ 注意:该节点目前处于开发/调试阶段,输出结果可能异常,暂不建议生产环境使用。
|
||||
|
||||
- **功能**:将图片背景转换为透明,支持多种细节处理。
|
||||
- **参数**:
|
||||
- `image`:输入图片。
|
||||
@@ -42,12 +52,12 @@
|
||||
- `mask`:前景掩码
|
||||
|
||||
## 典型用法
|
||||
1. 用 `LoadBiRefNetModel_RBS` 或 `LoadBiRefNetModelV2_RBS` 加载模型。
|
||||
2. 用 `BiRefNetUltraV2_RBS` 进行背景移除。
|
||||
1. 用 `LoadBiRefNetModelV3_RBS` 加载模型(推荐选择 BiRefNet_dynamic)。
|
||||
2. 用 `BiRefNetUltraV3_RBS` 进行背景移除。
|
||||
3. 可选:用 `TransparentBackgroundUltra_RBS` 进一步处理透明背景。
|
||||
|
||||
## 注意事项
|
||||
- 请将模型文件放在 `ComfyUI/models/BiRefNet/pth/` 目录下,或使用新版节点自动下载。
|
||||
- 请将模型文件放在 `ComfyUI/models/transparent-background/` 目录下,或使用新版节点自动下载。
|
||||
- 推荐使用 CUDA 设备以获得更快推理速度。
|
||||
- 细节处理方法对边缘质量有显著影响,可根据实际需求调整。
|
||||
- 插件所有节点均归类于 `RemoveBackgroundSuite`,便于统一管理。
|
||||
@@ -63,4 +73,8 @@ pip install -r requirements.txt
|
||||
- **节点不显示**:请确认插件已放入 `custom_nodes` 目录并重启 ComfyUI。
|
||||
|
||||
---
|
||||
|
||||
## 致谢
|
||||
本插件大量借鉴和参考了 [ComfyUI_LayerStyle_Advance](https://github.com/chflame163/ComfyUI_LayerStyle_Advance) 项目的设计与实现,特别感谢原作者 chflame163 的开源贡献!
|
||||
|
||||
如有更多问题请参考原项目文档或在 Issues 区反馈。
|
||||
+65
-22
@@ -11,9 +11,12 @@ import torch.nn.functional as F
|
||||
from torchvision import transforms
|
||||
from transformers import AutoModelForImageSegmentation
|
||||
import sys
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'BiRefNet_v2'))
|
||||
import math
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'BiRefNet'))
|
||||
|
||||
def get_files(path, extensions):
|
||||
if isinstance(extensions, list):
|
||||
extensions = tuple(extensions)
|
||||
files = {}
|
||||
for file in os.listdir(path):
|
||||
if file.endswith(extensions):
|
||||
@@ -21,7 +24,7 @@ def get_files(path, extensions):
|
||||
return files
|
||||
|
||||
def scan_model():
|
||||
model_path = os.path.join(folder_paths.models_dir, 'BiRefNet')
|
||||
model_path = os.path.join(folder_paths.models_dir, 'transparent-background')
|
||||
model_ext = [".pth"]
|
||||
model_dict = get_files(model_path, model_ext)
|
||||
return model_dict
|
||||
@@ -74,29 +77,69 @@ def generate_VITMatte_trimap(mask, erode, dilate):
|
||||
trimap[dilate_mask > 0.5] = 0.5
|
||||
return Image.fromarray(trimap)
|
||||
|
||||
def generate_VITMatte(image, trimap, local_files_only=False, device='cuda', max_megapixels=2.0):
|
||||
from transformers import AutoModelForImageSegmentation
|
||||
model_path = os.path.join(folder_paths.models_dir, 'BiRefNet', 'VITMatte')
|
||||
def check_and_download_model(model_path, repo_id):
|
||||
model_path = os.path.join(folder_paths.models_dir, model_path)
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(model_path, exist_ok=True)
|
||||
print(f"Downloading {repo_id} model...")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="ZhengPeng7/VITMatte", local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
|
||||
model = AutoModelForImageSegmentation.from_pretrained(model_path, trust_remote_code=True)
|
||||
model.to(device)
|
||||
model.eval()
|
||||
image = np.array(image)
|
||||
trimap = np.array(trimap)
|
||||
image = cv2.resize(image, (1024, 1024))
|
||||
trimap = cv2.resize(trimap, (1024, 1024))
|
||||
image = torch.from_numpy(image).permute(2, 0, 1).unsqueeze(0).float() / 255.0
|
||||
trimap = torch.from_numpy(trimap).unsqueeze(0).unsqueeze(0).float()
|
||||
image = image.to(device)
|
||||
trimap = trimap.to(device)
|
||||
snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt", "onnx", ".git"])
|
||||
return model_path
|
||||
|
||||
class VITMatteModel:
|
||||
def __init__(self,model,processor):
|
||||
self.model = model
|
||||
self.processor = processor
|
||||
|
||||
def load_VITMatte_model(model_name:str, local_files_only:bool=False) -> object:
|
||||
model_name = "vitmatte"
|
||||
model_repo = "hustvl/vitmatte-small-composition-1k"
|
||||
model_path = check_and_download_model(model_name, model_repo)
|
||||
from transformers import VitMatteImageProcessor, VitMatteForImageMatting
|
||||
model = VitMatteForImageMatting.from_pretrained(model_path, local_files_only=local_files_only)
|
||||
processor = VitMatteImageProcessor.from_pretrained(model_path, local_files_only=local_files_only)
|
||||
vitmatte = VITMatteModel(model, processor)
|
||||
return vitmatte
|
||||
|
||||
def generate_VITMatte(image, trimap, local_files_only=False, device='cuda', max_megapixels=2.0):
|
||||
import torch
|
||||
from PIL import Image
|
||||
if image.mode != 'RGB':
|
||||
image = image.convert('RGB')
|
||||
if trimap.mode != 'L':
|
||||
trimap = trimap.convert('L')
|
||||
max_megapixels *= 1048576
|
||||
width, height = image.size
|
||||
ratio = width / height
|
||||
target_width = math.sqrt(ratio * max_megapixels)
|
||||
target_height = target_width / ratio
|
||||
target_width = int(target_width)
|
||||
target_height = int(target_height)
|
||||
if width * height > max_megapixels:
|
||||
image = image.resize((target_width, target_height), Image.BILINEAR)
|
||||
trimap = trimap.resize((target_width, target_height), Image.BILINEAR)
|
||||
model_name = "hustvl/vitmatte-small-composition-1k"
|
||||
if device=="cpu":
|
||||
device = torch.device('cpu')
|
||||
else:
|
||||
if torch.cuda.is_available():
|
||||
device = torch.device('cuda')
|
||||
else:
|
||||
print("vitmatte device is set to cuda, but not available, using cpu instead.")
|
||||
device = torch.device('cpu')
|
||||
vit_matte_model = load_VITMatte_model(model_name=model_name, local_files_only=local_files_only)
|
||||
vit_matte_model.model.to(device)
|
||||
inputs = vit_matte_model.processor(images=image, trimaps=trimap, return_tensors="pt")
|
||||
with torch.no_grad():
|
||||
pred = model(image, trimap)
|
||||
pred = pred.cpu().numpy().squeeze()
|
||||
pred = cv2.resize(pred, (image.shape[3], image.shape[2]))
|
||||
return Image.fromarray((pred * 255).astype(np.uint8))
|
||||
inputs = {k: v.to(device) for k, v in inputs.items()}
|
||||
predictions = vit_matte_model.model(**inputs).alphas
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
mask = tensor2pil(predictions).convert('L')
|
||||
mask = mask.crop((0, 0, image.width, image.height))
|
||||
if width * height > max_megapixels:
|
||||
mask = mask.resize((width, height), Image.BILINEAR)
|
||||
return mask
|
||||
|
||||
def histogram_remap(mask, black_point, white_point):
|
||||
mask = tensor2pil(mask)
|
||||
|
||||
@@ -10,7 +10,9 @@ import tqdm
|
||||
from torchvision import transforms
|
||||
from transformers import AutoModelForImageSegmentation
|
||||
import sys
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'BiRefNet_v2'))
|
||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'BiRefNet'))
|
||||
from .BiRefNet.models.birefnet import BiRefNet
|
||||
from .BiRefNet.utils import check_state_dict
|
||||
|
||||
# 获取本地所有BiRefNet模型文件
|
||||
# 返回字典:{模型文件名: 路径}
|
||||
@@ -20,104 +22,154 @@ def get_models():
|
||||
model_dict = get_files(model_path, model_ext)
|
||||
return model_dict
|
||||
|
||||
# 加载本地BiRefNet模型节点
|
||||
class LoadBiRefNetModel_RBS:
|
||||
# 透明背景超强节点
|
||||
class TransparentBackgroundUltra_RBS:
|
||||
def __init__(self):
|
||||
self.birefnet = None
|
||||
self.state_dict = None
|
||||
self.NODE_NAME = 'TransparentBackgroundUltra_RBS'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# 自动扫描本地模型文件,优先显示推荐模型
|
||||
tmp_list = list(get_models().keys())
|
||||
model_list = []
|
||||
if 'BiRefNet-general-epoch_244.pth' in tmp_list:
|
||||
model_list.append('BiRefNet-general-epoch_244.pth')
|
||||
tmp_list.remove('BiRefNet-general-epoch_244.pth')
|
||||
model_list.extend(tmp_list)
|
||||
|
||||
def INPUT_TYPES(cls):
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda','cpu']
|
||||
def scan_transparent_models():
|
||||
import glob
|
||||
model_file_list = glob.glob(os.path.join(folder_paths.models_dir, "transparent-background") + '/*.pth')
|
||||
model_dict = {}
|
||||
for i in range(len(model_file_list)):
|
||||
_, __filename = os.path.split(model_file_list[i])
|
||||
model_dict[__filename] = model_file_list[i]
|
||||
return model_dict
|
||||
return {
|
||||
"required": {
|
||||
"model": (model_list,), # 选择模型文件
|
||||
"image": ("IMAGE",),
|
||||
"model": (list(scan_transparent_models().keys()),),
|
||||
"detail_method": (method_list,),
|
||||
"detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
|
||||
"detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
|
||||
"black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
|
||||
"white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
|
||||
"process_detail": ("BOOLEAN", {"default": True}),
|
||||
"device": (device_list,),
|
||||
"max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
|
||||
},
|
||||
"optional": {
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BIREFNET_MODEL",)
|
||||
RETURN_NAMES = ("birefnet_model",)
|
||||
FUNCTION = "load_birefnet_model"
|
||||
RETURN_TYPES = ("IMAGE", "MASK", )
|
||||
RETURN_NAMES = ("image", "mask", )
|
||||
FUNCTION = "transparent_background_ultra"
|
||||
CATEGORY = 'RemoveBackgroundSuite'
|
||||
|
||||
def transparent_background_ultra(self, image, model, detail_method, detail_erode, detail_dilate,
|
||||
black_point, white_point, process_detail, device, max_megapixels):
|
||||
import glob
|
||||
from transparent_background import Remover
|
||||
mode_dict = {"ckpt_base.pth": "base", "ckpt_base_nightly.pth": "base-nightly", "ckpt_fast.pth": "fast"}
|
||||
ret_images = []
|
||||
ret_masks = []
|
||||
if detail_method == 'VITMatte(local)':
|
||||
local_files_only = True
|
||||
else:
|
||||
local_files_only = False
|
||||
model_file_list = glob.glob(os.path.join(folder_paths.models_dir, "transparent-background") + '/*.pth')
|
||||
model_dict = {}
|
||||
for i in range(len(model_file_list)):
|
||||
_, __filename = os.path.split(model_file_list[i])
|
||||
model_dict[__filename] = model_file_list[i]
|
||||
try:
|
||||
mode = mode_dict[model]
|
||||
except:
|
||||
mode = "base"
|
||||
remover = Remover(mode=mode, jit=False, device=device, ckpt=model_dict[model])
|
||||
for i in image:
|
||||
i = torch.unsqueeze(i, 0)
|
||||
orig_image = tensor2pil(i).convert('RGB')
|
||||
ret_image = remover.process(orig_image, type='rgba')
|
||||
_mask = ret_image.split()[3]
|
||||
_mask = adjust_levels(_mask, 64, 192)
|
||||
if process_detail:
|
||||
detail_range = detail_erode + detail_dilate
|
||||
_mask = pil2tensor(_mask)
|
||||
if detail_method == 'GuidedFilter':
|
||||
_mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
|
||||
_mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
|
||||
elif detail_method == 'PyMatting':
|
||||
_mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
|
||||
else:
|
||||
_trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
|
||||
_mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels)
|
||||
_mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
|
||||
ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
|
||||
ret_images.append(pil2tensor(ret_image))
|
||||
ret_masks.append(image2mask(_mask))
|
||||
log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
|
||||
return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
|
||||
|
||||
# 加载模型权重并返回模型对象
|
||||
def load_birefnet_model(self, model):
|
||||
from .BiRefNet_v2.models.birefnet import BiRefNet
|
||||
from .BiRefNet_v2.utils import check_state_dict
|
||||
model_dict = get_models()
|
||||
self.birefnet = BiRefNet(bb_pretrained=False)
|
||||
self.state_dict = torch.load(model_dict[model], map_location='cpu', weights_only=True)
|
||||
self.state_dict = check_state_dict(self.state_dict)
|
||||
self.birefnet.load_state_dict(self.state_dict)
|
||||
return (self.birefnet,)
|
||||
# ======================== V3 新增:支持 BiRefNet-Dynamic ========================
|
||||
|
||||
# 自动下载并加载BiRefNet新版模型节点
|
||||
class LoadBiRefNetModelV2_RBS:
|
||||
class LoadBiRefNetModelV3_RBS:
|
||||
def __init__(self):
|
||||
self.model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# 支持的模型版本列表
|
||||
# 新增 dynamic 选项
|
||||
model_list = list(s.birefnet_model_repos.keys())
|
||||
return {
|
||||
"required": {
|
||||
"version": (model_list,{"default": model_list[0]}), # 选择模型版本
|
||||
"version": (model_list, {"default": model_list[0]}), # 选择模型版本
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("BIREFNET_MODEL",)
|
||||
RETURN_NAMES = ("birefnet_model",)
|
||||
FUNCTION = "load_birefnet_model"
|
||||
FUNCTION = "load_birefnet_model_v3"
|
||||
CATEGORY = 'RemoveBackgroundSuite'
|
||||
|
||||
# Huggingface仓库映射
|
||||
# 支持 dynamic、HR、HR-matting 模型
|
||||
birefnet_model_repos = {
|
||||
"BiRefNet-General": "ZhengPeng7/BiRefNet",
|
||||
"RMBG-2.0": "briaai/RMBG-2.0"
|
||||
"RMBG-2.0": "briaai/RMBG-2.0",
|
||||
"BiRefNet_dynamic": "ZhengPeng7/BiRefNet_dynamic",
|
||||
"BiRefNet_HR": "ZhengPeng7/BiRefNet_HR",
|
||||
"BiRefNet_HR-matting": "ZhengPeng7/BiRefNet_HR-matting"
|
||||
}
|
||||
|
||||
# 自动下载并加载模型
|
||||
def load_birefnet_model(self, version):
|
||||
def load_birefnet_model_v3(self, version):
|
||||
birefnet_path = os.path.join(folder_paths.models_dir, 'BiRefNet')
|
||||
os.makedirs(birefnet_path, exist_ok=True)
|
||||
|
||||
model_path = os.path.join(birefnet_path, version)
|
||||
|
||||
# 兼容老模型
|
||||
if version == "BiRefNet-General":
|
||||
old_birefnet_path = os.path.join(birefnet_path, 'pth')
|
||||
old_model = "BiRefNet-general-epoch_244.pth"
|
||||
old_model_path = os.path.join(old_birefnet_path, old_model)
|
||||
if os.path.exists(old_model_path):
|
||||
from .BiRefNet_v2.models.birefnet import BiRefNet
|
||||
from .BiRefNet_v2.utils import check_state_dict
|
||||
from .BiRefNet.models.birefnet import BiRefNet
|
||||
from .BiRefNet.utils import check_state_dict
|
||||
self.birefnet = BiRefNet(bb_pretrained=False)
|
||||
self.state_dict = torch.load(old_model_path, map_location='cpu', weights_only=True)
|
||||
self.state_dict = check_state_dict(self.state_dict)
|
||||
self.birefnet.load_state_dict(self.state_dict)
|
||||
return (self.birefnet,)
|
||||
# 若本地无模型则自动下载
|
||||
elif not os.path.exists(model_path):
|
||||
# 动态模型及HR、HR-matting模型下载
|
||||
if version in ["BiRefNet_dynamic", "BiRefNet_HR", "BiRefNet_HR-matting"] and not os.path.exists(model_path):
|
||||
log(f"Downloading {version} model...")
|
||||
repo_id = self.birefnet_model_repos[version]
|
||||
from huggingface_hub import snapshot_download
|
||||
repo_id = self.birefnet_model_repos[version]
|
||||
snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
|
||||
elif version == "RMBG-2.0" and not os.path.exists(model_path):
|
||||
log(f"Downloading RMBG-2.0 model...")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="briaai/RMBG-2.0", local_dir=model_path, ignore_patterns=["*.md", "*.txt"])
|
||||
|
||||
self.model = AutoModelForImageSegmentation.from_pretrained(model_path, trust_remote_code=True)
|
||||
return (self.model,)
|
||||
|
||||
# BiRefNet Ultra V2 背景移除主节点
|
||||
class BiRefNetUltraV2_RBS:
|
||||
class BiRefNetUltraV3_RBS:
|
||||
def __init__(self):
|
||||
self.NODE_NAME = 'BiRefNetUltraV2_RBS'
|
||||
self.NODE_NAME = 'BiRefNetUltraV3_RBS'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -143,11 +195,11 @@ class BiRefNetUltraV2_RBS:
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", )
|
||||
RETURN_NAMES = ("image", "mask", )
|
||||
FUNCTION = "birefnet_ultra_v2"
|
||||
FUNCTION = "birefnet_ultra_v3"
|
||||
CATEGORY = 'RemoveBackgroundSuite'
|
||||
|
||||
# 主推理流程
|
||||
def birefnet_ultra_v2(self, image, birefnet_model, detail_method, detail_erode, detail_dilate,
|
||||
# 主推理流程,兼容 dynamic 模型
|
||||
def birefnet_ultra_v3(self, image, birefnet_model, detail_method, detail_erode, detail_dilate,
|
||||
black_point, white_point, process_detail, device, max_megapixels):
|
||||
ret_images = []
|
||||
ret_masks = []
|
||||
@@ -162,7 +214,7 @@ class BiRefNetUltraV2_RBS:
|
||||
birefnet_model.eval()
|
||||
|
||||
comfy_pbar = ProgressBar(len(image))
|
||||
tqdm_pbar = tqdm.tqdm(total=len(image), desc="Processing BiRefNet")
|
||||
tqdm_pbar = tqdm.tqdm(total=len(image), desc="Processing BiRefNetV3")
|
||||
for i in image:
|
||||
i = torch.unsqueeze(i, 0)
|
||||
orig_image = tensor2pil(i).convert('RGB')
|
||||
@@ -214,95 +266,15 @@ class BiRefNetUltraV2_RBS:
|
||||
log(f"{self.NODE_NAME} Processed {len(ret_masks)} image(s).", message_type='finish')
|
||||
return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
|
||||
|
||||
# 透明背景超强节点
|
||||
class TransparentBackgroundUltra_RBS:
|
||||
def __init__(self):
|
||||
self.NODE_NAME = 'TransparentBackgroundUltra_RBS'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda','cpu']
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",), # 输入图片
|
||||
"model": (list(scan_model().keys()),), # 选择模型
|
||||
"detail_method": (method_list,), # 细节处理方式
|
||||
"detail_erode": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
|
||||
"detail_dilate": ("INT", {"default": 6, "min": 1, "max": 255, "step": 1}),
|
||||
"black_point": ("FLOAT", {"default": 0.01, "min": 0.01, "max": 0.98, "step": 0.01, "display": "slider"}),
|
||||
"white_point": ("FLOAT", {"default": 0.99, "min": 0.02, "max": 0.99, "step": 0.01, "display": "slider"}),
|
||||
"process_detail": ("BOOLEAN", {"default": True}),
|
||||
"device": (device_list,),
|
||||
"max_megapixels": ("FLOAT", {"default": 2.0, "min": 1, "max": 999, "step": 0.1}),
|
||||
},
|
||||
"optional": {
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", )
|
||||
RETURN_NAMES = ("image", "mask", )
|
||||
FUNCTION = "transparent_background_ultra"
|
||||
CATEGORY = 'RemoveBackgroundSuite'
|
||||
|
||||
# 主推理流程
|
||||
def transparent_background_ultra(self, image, model, detail_method, detail_erode, detail_dilate,
|
||||
black_point, white_point, process_detail, device, max_megapixels):
|
||||
|
||||
from transparent_background import Remover
|
||||
|
||||
ret_images = []
|
||||
ret_masks = []
|
||||
if detail_method == 'VITMatte(local)':
|
||||
local_files_only = True
|
||||
else:
|
||||
local_files_only = False
|
||||
model_dict = scan_model()
|
||||
try :
|
||||
mode = mode_dict[model]
|
||||
except :
|
||||
mode = "base"
|
||||
remover = Remover(mode=mode, jit=False, device=device, ckpt=model_dict[model])
|
||||
for i in image:
|
||||
i = torch.unsqueeze(i, 0)
|
||||
orig_image = tensor2pil(i).convert('RGB')
|
||||
ret_image = remover.process(orig_image, type='rgba')
|
||||
_mask = ret_image.split()[3]
|
||||
_mask = adjust_levels(_mask, 64, 192)
|
||||
|
||||
if process_detail:
|
||||
detail_range = detail_erode + detail_dilate
|
||||
_mask = pil2tensor(_mask)
|
||||
if detail_method == 'GuidedFilter':
|
||||
_mask = guided_filter_alpha(i, _mask, detail_range // 6 + 1)
|
||||
_mask = tensor2pil(histogram_remap(_mask, black_point, white_point))
|
||||
elif detail_method == 'PyMatting':
|
||||
_mask = tensor2pil(mask_edge_detail(i, _mask, detail_range // 8 + 1, black_point, white_point))
|
||||
else:
|
||||
_trimap = generate_VITMatte_trimap(_mask, detail_erode, detail_dilate)
|
||||
_mask = generate_VITMatte(orig_image, _trimap, local_files_only=local_files_only, device=device, max_megapixels=max_megapixels)
|
||||
_mask = tensor2pil(histogram_remap(pil2tensor(_mask), black_point, white_point))
|
||||
ret_image = RGB2RGBA(orig_image, _mask.convert('L'))
|
||||
|
||||
ret_images.append(pil2tensor(ret_image))
|
||||
ret_masks.append(image2mask(_mask))
|
||||
|
||||
log(f"{self.NODE_NAME} Processed {len(ret_images)} image(s).", message_type='finish')
|
||||
|
||||
return (torch.cat(ret_images, dim=0), torch.cat(ret_masks, dim=0),)
|
||||
|
||||
# 节点注册映射
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LoadBiRefNetModel_RBS": LoadBiRefNetModel_RBS,
|
||||
"LoadBiRefNetModelV2_RBS": LoadBiRefNetModelV2_RBS,
|
||||
"BiRefNetUltraV2_RBS": BiRefNetUltraV2_RBS,
|
||||
"TransparentBackgroundUltra_RBS": TransparentBackgroundUltra_RBS
|
||||
"TransparentBackgroundUltra_RBS": TransparentBackgroundUltra_RBS,
|
||||
"LoadBiRefNetModelV3_RBS": LoadBiRefNetModelV3_RBS,
|
||||
"BiRefNetUltraV3_RBS": BiRefNetUltraV3_RBS
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LoadBiRefNetModel_RBS": "Load BiRefNet Model (RBS)",
|
||||
"LoadBiRefNetModelV2_RBS": "Load BiRefNet Model V2 (RBS)",
|
||||
"BiRefNetUltraV2_RBS": "BiRefNet Ultra V2 (RBS)",
|
||||
"TransparentBackgroundUltra_RBS": "Transparent Background Ultra (RBS)"
|
||||
"TransparentBackgroundUltra_RBS": "Transparent Background Ultra (RBS) [WIP]",
|
||||
"LoadBiRefNetModelV3_RBS": "Load BiRefNet Model V3 (RBS)",
|
||||
"BiRefNetUltraV3_RBS": "BiRefNet Ultra V3 (RBS)"
|
||||
}
|
||||
Reference in New Issue
Block a user