diff --git a/BiRefNet_node_config.py b/BiRefNet_node_config.py index 6d6dc09..4b4f600 100644 --- a/BiRefNet_node_config.py +++ b/BiRefNet_node_config.py @@ -2,7 +2,7 @@ import os import math -# from folder_paths import models_dir + os.environ['HOME'] = os.path.expanduser("~") class Config(): def __init__(self) -> None: diff --git a/__init__.py b/__init__.py index 604cbbc..d721463 100644 --- a/__init__.py +++ b/__init__.py @@ -1,10 +1,3 @@ -import os -import folder_paths - -rembg_path = os.path.join(folder_paths.models_dir, 'rembg') -if not os.path.exists(rembg_path): - os.makedirs(rembg_path) - from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/assets/Colors_Presets.png b/assets/Colors_Presets.png deleted file mode 100644 index aef3799..0000000 Binary files a/assets/Colors_Presets.png and /dev/null differ diff --git a/nodes.py b/nodes.py index a0122d7..36fc804 100644 --- a/nodes.py +++ b/nodes.py @@ -2,15 +2,20 @@ from transformers import AutoModelForImageSegmentation import torch from torchvision import transforms import numpy as np -from PIL import Image, ImageColor +from PIL import Image import torch.nn.functional as F from .BiRefNet_node_config import Config -import folder_paths -import os import comfy.model_management as mm -from huggingface_hub import snapshot_download -comfyui_models_dir = folder_paths.models_dir +Config() + +torch.set_float32_matmul_precision(["high", "highest"][0]) + +birefnet = AutoModelForImageSegmentation.from_pretrained( + "ZhengPeng7/BiRefNet", trust_remote_code=True +) + + transform_image = transforms.Compose( [ @@ -43,84 +48,28 @@ def get_device_by_name(device): device = "cpu" if torch.cuda.is_available(): device = "cuda" - # device = torch.device("cuda") elif torch.backends.mps.is_available(): device = "mps" - # device = torch.device("mps") elif torch.xpu.is_available(): device = "xpu" - # device = torch.device("xpu") - # device = mm.get_torch_device() except: raise AttributeError("What's your device(到底用什么设备跑的)?") - # elif device == 'cuda': - # device = torch.device("cuda") - # elif device == "mps": - # device = torch.device("mps") - # elif device == "xpu": - # device = torch.device("xpu") print("\033[93mUse Device(使用设备):", device, "\033[0m") return device -def get_dtype_by_name(dtype): - """ - "dtype": (["auto","fp16","bf16","fp32", "fp8_e4m3fn", "fp8_e4m3fnuz", "fp8_e5m2", "fp8_e5m2fnuz"],{"default":"auto"}), - """ - if dtype == 'auto': - try: - if mm.should_use_fp16(): - dtype = torch.float16 - elif mm.should_use_bf16(): - dtype = torch.bfloat16 - else: - dtype = torch.float32 - except: - raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.") - elif dtype== "fp16": - dtype = torch.float16 - elif dtype == "bf16": - dtype = torch.bfloat16 - elif dtype == "fp32": - dtype = torch.float32 - elif dtype == "fp8_e4m3fn": - dtype = torch.float8_e4m3fn - elif dtype == "fp8_e4m3fnuz": - dtype = torch.float8_e4m3fnuz - elif dtype == "fp8_e5m2": - dtype = torch.float8_e5m2 - elif dtype == "fp8_e5m2fnuz": - dtype = torch.float8_e5m2fnuz - print("\033[93mModel Precision(模型精度):", dtype, "\033[0m") - return dtype - class BiRefNet_Hugo: def __init__(self): - self.model = None - self.loaded_model_name = None + pass @classmethod def INPUT_TYPES(cls): - rembg_list = os.listdir(os.path.join(comfyui_models_dir, "rembg")) - rembg_list.insert(0, "Auto_DownLoad-ZhengPeng7/BiRefNet") - rembg_list.insert(1, "Auto_DownLoad-ZhengPeng7/BiRefNet-DIS5K-TR_TEs") - rembg_list.insert(2, "Auto_DownLoad-ZhengPeng7/BiRefNet-COD") - rembg_list.insert(3, "Auto_DownLoad-ZhengPeng7/BiRefNet-HRSOD") - rembg_list.insert(4, "Auto_DownLoad-ZhengPeng7/BiRefNet-portrait") + return { "required": { - "model": (rembg_list, ), - # "model": (["Auto_Download"] + os.listdir(os.path.join(comfyui_models_dir, "rembg"))), "image": ("IMAGE",), - # "background_color_name": (["transparency", "green", "white", "red", "yellow", "blue", "black", "pink", "purple", "brown"],{"default": "transparency"}), "background_color_name": (colors,{"default": "transparency"}), - "background_color_code": ("STRING",{"default": "00ffdd"}), - "background_color_mode": ("BOOLEAN", {"default": True, "label_on": "color_name", "label_off": "color_code"}), - "device": (["auto", "cuda", "cpu", "mps", "xpu", "meta"],{"default": "auto"}), - "dtype": (["auto","fp16","bf16","fp32", "fp8_e4m3fn", "fp8_e4m3fnuz", "fp8_e5m2", "fp8_e5m2fnuz"],{"default":"fp32"}), - "cpu_offload": ("BOOLEAN", {"default": False, "label_on": "model_to_cpu", "label_off": "unload_model"}), - "Auto_Download_Path": ("BOOLEAN", {"default": True, "label_on": "rembg_local本地", "label_off": ".cache缓存"}), - "Show_Colors_In_Ternimal": ("BOOLEAN", {"default": False, "label_on": "yes", "label_off": "no"}), + "device": (["auto", "cuda", "cpu", "mps", "xpu", "meta"],{"default": "auto"}) } } @@ -131,73 +80,34 @@ class BiRefNet_Hugo: def background_remove(self, image, - model, device, - dtype, - cpu_offload, background_color_name, - background_color_code, - background_color_mode, - Auto_Download_Path, - Show_Colors_In_Ternimal, ): - if Show_Colors_In_Ternimal: - for name, code in ImageColor.colormap.items(): - print( f'{name:30} : {code}' ) - Config() - torch.set_float32_matmul_precision(["high", "highest"][1]) processed_images = [] processed_masks = [] - model_name = model.replace("Auto_DownLoad-", "") - if 'Auto_DownLoad-' not in model: - model_path = os.path.join(comfyui_models_dir, "rembg", model) - elif ('Auto_DownLoad-' in model) and (Auto_Download_Path == True): - model_path = os.path.join(comfyui_models_dir, "rembg", ("models--" + str(model_name).replace("/", "--"))) - if not os.path.exists(os.path.join(model_path, "model.safetensors")): - snapshot_download(model_name, - local_dir=model_path, - local_dir_use_symlinks=False - ) - elif ('Auto_DownLoad-' in model) and (Auto_Download_Path == False): - model_path = model_name - + device = get_device_by_name(device) - dtype = get_dtype_by_name(dtype) - - if self.loaded_model_name != model: - del self.model - self.model = None - if self.model == None: - self. model = AutoModelForImageSegmentation.from_pretrained( - model_path, - trust_remote_code=True, - ).to(device, dtype) - else: - self.model.to(device) - self.model.eval() + birefnet.to(device) for image in image: orig_image = tensor2pil(image) w,h = orig_image.size image = resize_image(orig_image) im_tensor = transform_image(image).unsqueeze(0) - im_tensor=im_tensor.to(device, dtype) + im_tensor=im_tensor.to(device) with torch.no_grad(): - result = self.model(im_tensor)[-1].sigmoid().cpu() + result = birefnet(im_tensor)[-1].sigmoid().cpu() result = torch.squeeze(F.interpolate(result, size=(h,w))) ma = torch.max(result) mi = torch.min(result) result = (result-mi)/(ma-mi) im_array = (result*255).cpu().data.numpy().astype(np.uint8) pil_im = Image.fromarray(np.squeeze(im_array)) - if background_color_name == 'transparency' and background_color_mode == True: + if background_color_name == 'transparency': color = (0,0,0,0) mode = "RGBA" else: color = background_color_name mode = "RGB" - if not background_color_mode: - color = "#" + str(background_color_code).replace("#", "").replace(":", "").replace(" ", "") # 颜色码-绿色:#00FF00 - # new_im = Image.new("RGBA", pil_im.size, (0,0,0,0)) new_im = Image.new(mode, pil_im.size, color) new_im.paste(orig_image, mask=pil_im) new_im_tensor = pil2tensor(new_im) @@ -207,12 +117,6 @@ class BiRefNet_Hugo: new_ims = torch.cat(processed_images, dim=0) new_masks = torch.cat(processed_masks, dim=0) - if cpu_offload == True: - self.model.to("cpu") - self.loaded_model_name = model - else: - del self.model - self.model = None return new_ims, new_masks