From c0af2876cf2576ccd1f173782a7ab1c1f71cf6b1 Mon Sep 17 00:00:00 2001 From: Cyber Dick Lang <286878701@qq.com> Date: Sun, 1 Jun 2025 02:10:04 +0800 Subject: [PATCH] =?UTF-8?q?v1.1:=20=E4=BC=98=E5=8C=96BiRefNet=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E5=8A=A0=E8=BD=BD=E5=92=8C=E6=8E=A8=E7=90=86=E6=B5=81?= =?UTF-8?q?=E7=A8=8B=EF=BC=8C=E6=94=B9=E8=BF=9B=E6=96=87=E6=A1=A3=E7=BB=93?= =?UTF-8?q?=E6=9E=84=EF=BC=8C=E7=A7=BB=E9=99=A4=E7=89=88=E6=9C=AC=E5=8F=B7?= =?UTF-8?q?=E4=BE=9D=E8=B5=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- {BiRefNet_v2 => BiRefNet}/LICENSE | 0 {BiRefNet_v2 => BiRefNet}/README.md | 0 {BiRefNet_v2 => BiRefNet}/__init__.py | 0 {BiRefNet_v2 => BiRefNet}/config.py | 0 {BiRefNet_v2 => BiRefNet}/dataset.py | 0 .../eval_existingOnes.py | 0 .../evaluation/metrics.py | 0 {BiRefNet_v2 => BiRefNet}/gen_best_ep.py | 0 {BiRefNet_v2 => BiRefNet}/image_proc.py | 0 {BiRefNet_v2 => BiRefNet}/inference.py | 0 {BiRefNet_v2 => BiRefNet}/loss.py | 0 {BiRefNet_v2 => BiRefNet}/make_a_copy.sh | 0 .../models/backbones/build_backbone.py | 0 .../models/backbones/pvt_v2.py | 0 .../models/backbones/swin_v1.py | 0 {BiRefNet_v2 => BiRefNet}/models/birefnet.py | 0 .../models/modules/aspp.py | 0 .../models/modules/decoder_blocks.py | 0 .../models/modules/deform_conv.py | 0 .../models/modules/lateral_blocks.py | 0 .../models/modules/mlp.py | 0 .../models/modules/prompt_encoder.py | 0 .../models/modules/utils.py | 0 .../models/refinement/refiner.py | 0 .../models/refinement/stem_layer.py | 0 {BiRefNet_v2 => BiRefNet}/requirements.txt | 0 {BiRefNet_v2 => BiRefNet}/rm_cache.sh | 0 {BiRefNet_v2 => BiRefNet}/sub.sh | 0 {BiRefNet_v2 => BiRefNet}/test.sh | 0 {BiRefNet_v2 => BiRefNet}/train.py | 0 {BiRefNet_v2 => BiRefNet}/train.sh | 0 {BiRefNet_v2 => BiRefNet}/train_test.sh | 0 .../tutorials/BiRefNet_inference.ipynb | 0 .../tutorials/BiRefNet_pth2onnx.ipynb | 0 {BiRefNet_v2 => BiRefNet}/utils.py | 0 README.md | 60 +++-- imagefunc.py | 87 ++++-- nodes.py | 248 ++++++++---------- 38 files changed, 212 insertions(+), 183 deletions(-) rename {BiRefNet_v2 => BiRefNet}/LICENSE (100%) rename {BiRefNet_v2 => BiRefNet}/README.md (100%) rename {BiRefNet_v2 => BiRefNet}/__init__.py (100%) rename {BiRefNet_v2 => BiRefNet}/config.py (100%) rename {BiRefNet_v2 => BiRefNet}/dataset.py (100%) rename {BiRefNet_v2 => BiRefNet}/eval_existingOnes.py (100%) rename {BiRefNet_v2 => BiRefNet}/evaluation/metrics.py (100%) rename {BiRefNet_v2 => BiRefNet}/gen_best_ep.py (100%) rename {BiRefNet_v2 => BiRefNet}/image_proc.py (100%) rename {BiRefNet_v2 => BiRefNet}/inference.py (100%) rename {BiRefNet_v2 => BiRefNet}/loss.py (100%) rename {BiRefNet_v2 => BiRefNet}/make_a_copy.sh (100%) rename {BiRefNet_v2 => BiRefNet}/models/backbones/build_backbone.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/backbones/pvt_v2.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/backbones/swin_v1.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/birefnet.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/aspp.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/decoder_blocks.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/deform_conv.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/lateral_blocks.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/mlp.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/prompt_encoder.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/modules/utils.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/refinement/refiner.py (100%) rename {BiRefNet_v2 => BiRefNet}/models/refinement/stem_layer.py (100%) rename {BiRefNet_v2 => BiRefNet}/requirements.txt (100%) rename {BiRefNet_v2 => BiRefNet}/rm_cache.sh (100%) rename {BiRefNet_v2 => BiRefNet}/sub.sh (100%) rename {BiRefNet_v2 => BiRefNet}/test.sh (100%) rename {BiRefNet_v2 => BiRefNet}/train.py (100%) rename {BiRefNet_v2 => BiRefNet}/train.sh (100%) rename {BiRefNet_v2 => BiRefNet}/train_test.sh (100%) rename {BiRefNet_v2 => BiRefNet}/tutorials/BiRefNet_inference.ipynb (100%) rename {BiRefNet_v2 => BiRefNet}/tutorials/BiRefNet_pth2onnx.ipynb (100%) rename {BiRefNet_v2 => BiRefNet}/utils.py (100%) diff --git a/BiRefNet_v2/LICENSE b/BiRefNet/LICENSE similarity index 100% rename from BiRefNet_v2/LICENSE rename to BiRefNet/LICENSE diff --git a/BiRefNet_v2/README.md b/BiRefNet/README.md similarity index 100% rename from BiRefNet_v2/README.md rename to BiRefNet/README.md diff --git a/BiRefNet_v2/__init__.py b/BiRefNet/__init__.py similarity index 100% rename from BiRefNet_v2/__init__.py rename to BiRefNet/__init__.py diff --git a/BiRefNet_v2/config.py b/BiRefNet/config.py similarity index 100% rename from BiRefNet_v2/config.py rename to BiRefNet/config.py diff --git a/BiRefNet_v2/dataset.py b/BiRefNet/dataset.py similarity index 100% rename from BiRefNet_v2/dataset.py rename to BiRefNet/dataset.py diff --git a/BiRefNet_v2/eval_existingOnes.py b/BiRefNet/eval_existingOnes.py similarity index 100% rename from BiRefNet_v2/eval_existingOnes.py rename to BiRefNet/eval_existingOnes.py diff --git a/BiRefNet_v2/evaluation/metrics.py b/BiRefNet/evaluation/metrics.py similarity index 100% rename from BiRefNet_v2/evaluation/metrics.py rename to BiRefNet/evaluation/metrics.py diff --git a/BiRefNet_v2/gen_best_ep.py b/BiRefNet/gen_best_ep.py similarity index 100% rename from BiRefNet_v2/gen_best_ep.py rename to BiRefNet/gen_best_ep.py diff --git a/BiRefNet_v2/image_proc.py b/BiRefNet/image_proc.py similarity index 100% rename from BiRefNet_v2/image_proc.py rename to BiRefNet/image_proc.py diff --git a/BiRefNet_v2/inference.py b/BiRefNet/inference.py similarity index 100% rename from BiRefNet_v2/inference.py rename to BiRefNet/inference.py diff --git a/BiRefNet_v2/loss.py b/BiRefNet/loss.py similarity index 100% rename from BiRefNet_v2/loss.py rename to BiRefNet/loss.py diff --git a/BiRefNet_v2/make_a_copy.sh b/BiRefNet/make_a_copy.sh similarity index 100% rename from BiRefNet_v2/make_a_copy.sh rename to BiRefNet/make_a_copy.sh diff --git a/BiRefNet_v2/models/backbones/build_backbone.py b/BiRefNet/models/backbones/build_backbone.py similarity index 100% rename from BiRefNet_v2/models/backbones/build_backbone.py rename to BiRefNet/models/backbones/build_backbone.py diff --git a/BiRefNet_v2/models/backbones/pvt_v2.py b/BiRefNet/models/backbones/pvt_v2.py similarity index 100% rename from BiRefNet_v2/models/backbones/pvt_v2.py rename to BiRefNet/models/backbones/pvt_v2.py diff --git a/BiRefNet_v2/models/backbones/swin_v1.py b/BiRefNet/models/backbones/swin_v1.py similarity index 100% rename from BiRefNet_v2/models/backbones/swin_v1.py rename to BiRefNet/models/backbones/swin_v1.py diff --git a/BiRefNet_v2/models/birefnet.py b/BiRefNet/models/birefnet.py similarity index 100% rename from BiRefNet_v2/models/birefnet.py rename to BiRefNet/models/birefnet.py diff --git a/BiRefNet_v2/models/modules/aspp.py b/BiRefNet/models/modules/aspp.py similarity index 100% rename from BiRefNet_v2/models/modules/aspp.py rename to BiRefNet/models/modules/aspp.py diff --git a/BiRefNet_v2/models/modules/decoder_blocks.py b/BiRefNet/models/modules/decoder_blocks.py similarity index 100% rename from BiRefNet_v2/models/modules/decoder_blocks.py rename to BiRefNet/models/modules/decoder_blocks.py diff --git a/BiRefNet_v2/models/modules/deform_conv.py b/BiRefNet/models/modules/deform_conv.py similarity index 100% rename from BiRefNet_v2/models/modules/deform_conv.py rename to BiRefNet/models/modules/deform_conv.py diff --git a/BiRefNet_v2/models/modules/lateral_blocks.py b/BiRefNet/models/modules/lateral_blocks.py similarity index 100% rename from BiRefNet_v2/models/modules/lateral_blocks.py rename to BiRefNet/models/modules/lateral_blocks.py diff --git a/BiRefNet_v2/models/modules/mlp.py b/BiRefNet/models/modules/mlp.py similarity index 100% rename from BiRefNet_v2/models/modules/mlp.py rename to BiRefNet/models/modules/mlp.py diff --git a/BiRefNet_v2/models/modules/prompt_encoder.py b/BiRefNet/models/modules/prompt_encoder.py similarity index 100% rename from BiRefNet_v2/models/modules/prompt_encoder.py rename to BiRefNet/models/modules/prompt_encoder.py diff --git a/BiRefNet_v2/models/modules/utils.py b/BiRefNet/models/modules/utils.py similarity index 100% rename from BiRefNet_v2/models/modules/utils.py rename to BiRefNet/models/modules/utils.py diff --git a/BiRefNet_v2/models/refinement/refiner.py b/BiRefNet/models/refinement/refiner.py similarity index 100% rename from BiRefNet_v2/models/refinement/refiner.py rename to BiRefNet/models/refinement/refiner.py diff --git a/BiRefNet_v2/models/refinement/stem_layer.py b/BiRefNet/models/refinement/stem_layer.py similarity index 100% rename from BiRefNet_v2/models/refinement/stem_layer.py rename to BiRefNet/models/refinement/stem_layer.py diff --git a/BiRefNet_v2/requirements.txt b/BiRefNet/requirements.txt similarity index 100% rename from BiRefNet_v2/requirements.txt rename to BiRefNet/requirements.txt diff --git a/BiRefNet_v2/rm_cache.sh b/BiRefNet/rm_cache.sh similarity index 100% rename from BiRefNet_v2/rm_cache.sh rename to BiRefNet/rm_cache.sh diff --git a/BiRefNet_v2/sub.sh b/BiRefNet/sub.sh similarity index 100% rename from BiRefNet_v2/sub.sh rename to BiRefNet/sub.sh diff --git a/BiRefNet_v2/test.sh b/BiRefNet/test.sh similarity index 100% rename from BiRefNet_v2/test.sh rename to BiRefNet/test.sh diff --git a/BiRefNet_v2/train.py b/BiRefNet/train.py similarity index 100% rename from BiRefNet_v2/train.py rename to BiRefNet/train.py diff --git a/BiRefNet_v2/train.sh b/BiRefNet/train.sh similarity index 100% rename from BiRefNet_v2/train.sh rename to BiRefNet/train.sh diff --git a/BiRefNet_v2/train_test.sh b/BiRefNet/train_test.sh similarity index 100% rename from BiRefNet_v2/train_test.sh rename to BiRefNet/train_test.sh diff --git a/BiRefNet_v2/tutorials/BiRefNet_inference.ipynb b/BiRefNet/tutorials/BiRefNet_inference.ipynb similarity index 100% rename from BiRefNet_v2/tutorials/BiRefNet_inference.ipynb rename to BiRefNet/tutorials/BiRefNet_inference.ipynb diff --git a/BiRefNet_v2/tutorials/BiRefNet_pth2onnx.ipynb b/BiRefNet/tutorials/BiRefNet_pth2onnx.ipynb similarity index 100% rename from BiRefNet_v2/tutorials/BiRefNet_pth2onnx.ipynb rename to BiRefNet/tutorials/BiRefNet_pth2onnx.ipynb diff --git a/BiRefNet_v2/utils.py b/BiRefNet/utils.py similarity index 100% rename from BiRefNet_v2/utils.py rename to BiRefNet/utils.py diff --git a/README.md b/README.md index 8d30de9..99f969c 100644 --- a/README.md +++ b/README.md @@ -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 区反馈。 \ No newline at end of file diff --git a/imagefunc.py b/imagefunc.py index c6e8c24..2318581 100644 --- a/imagefunc.py +++ b/imagefunc.py @@ -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) diff --git a/nodes.py b/nodes.py index f93fa05..120738b 100644 --- a/nodes.py +++ b/nodes.py @@ -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)" } \ No newline at end of file