添加设备选择和背景颜色选择

This commit is contained in:
yueliang
2024-09-05 11:21:41 +08:00
parent f483313270
commit 37a20e516c
4 changed files with 19 additions and 122 deletions
+1 -1
View File
@@ -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:
-7
View File
@@ -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']
Binary file not shown.

Before

Width:  |  Height:  |  Size: 22 KiB

+18 -114
View File
@@ -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