Add files via upload
This commit is contained in:
+48
@@ -0,0 +1,48 @@
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import os
|
||||
import importlib.util
|
||||
|
||||
# Add module directory to Python path
|
||||
current_dir = Path(__file__).parent
|
||||
if str(current_dir) not in sys.path:
|
||||
sys.modules[__name__] = sys.modules.get(__name__, type(__name__, (), {}))
|
||||
sys.path.insert(0, str(current_dir))
|
||||
|
||||
# Initialize mappings
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
def load_nodes():
|
||||
"""Automatically discover and load node definitions"""
|
||||
for file in current_dir.glob("*.py"):
|
||||
if file.stem == "__init__":
|
||||
continue
|
||||
|
||||
try:
|
||||
# Import module
|
||||
spec = importlib.util.spec_from_file_location(file.stem, file)
|
||||
if spec and spec.loader:
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
# Update mappings
|
||||
if hasattr(module, "NODE_CLASS_MAPPINGS"):
|
||||
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
|
||||
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"):
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
# Initialize paths if available
|
||||
if hasattr(module, "Paths") and hasattr(module.Paths, "LLM_DIR"):
|
||||
os.makedirs(module.Paths.LLM_DIR, exist_ok=True)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error loading {file.name}: {str(e)}")
|
||||
|
||||
# Load all nodes
|
||||
load_nodes()
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS"
|
||||
]
|
||||
+132
@@ -0,0 +1,132 @@
|
||||
import os
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torchvision import transforms
|
||||
import numpy as np
|
||||
import folder_paths
|
||||
from transformers import AutoModelForImageSegmentation
|
||||
from PIL import ImageFilter
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
folder_paths.add_model_folder_path("rmbg", os.path.join(folder_paths.models_dir, "RMBG"))
|
||||
|
||||
AVAILABLE_MODELS = {
|
||||
"RMBG-2.0": "briaai/RMBG-2.0"
|
||||
}
|
||||
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
class AILAB_RMBG:
|
||||
def __init__(self):
|
||||
self.model = None
|
||||
self.current_model_version = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_version": (list(AVAILABLE_MODELS.keys()),),
|
||||
},
|
||||
"optional": {
|
||||
"sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"process_res": ("INT", {"default": 1024, "min": 512, "max": 2048, "step": 128}),
|
||||
"mask_blur": ("INT", {"default": 0, "min": 0, "max": 64, "step": 1}),
|
||||
"mask_offset": ("INT", {"default": 0, "min": -20, "max": 20, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "mask")
|
||||
FUNCTION = "remove_background"
|
||||
CATEGORY = "🧪AILab/🧽RMBG"
|
||||
|
||||
def remove_background(self, image, model_version, sensitivity=0.5, process_res=1024, mask_blur=0, mask_offset=0):
|
||||
try:
|
||||
cache_dir = os.path.join(folder_paths.models_dir, "RMBG", "RMBG-2.0")
|
||||
|
||||
if self.current_model_version != model_version or self.model is None:
|
||||
model_path = AVAILABLE_MODELS[model_version]
|
||||
model_files_path = os.path.join(cache_dir, model_version.replace("/", "--"))
|
||||
|
||||
if not os.path.exists(model_files_path):
|
||||
print(f"Downloading {model_version} model... This may take a while.")
|
||||
|
||||
self.model = AutoModelForImageSegmentation.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=True,
|
||||
cache_dir=cache_dir,
|
||||
revision="main",
|
||||
local_files_only=False
|
||||
)
|
||||
torch.set_float32_matmul_precision('high')
|
||||
self.model.to(device)
|
||||
self.model.eval()
|
||||
self.current_model_version = model_version
|
||||
print(f"Loaded model version: {model_version}")
|
||||
|
||||
transform_image = transforms.Compose([
|
||||
transforms.Resize((process_res, process_res)),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
|
||||
])
|
||||
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Error loading RMBG model: {str(e)}")
|
||||
|
||||
processed_images = []
|
||||
processed_masks = []
|
||||
|
||||
for img in image:
|
||||
orig_image = tensor2pil(img)
|
||||
input_tensor = transform_image(orig_image).unsqueeze(0).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
preds = self.model(input_tensor)[-1].sigmoid().cpu()
|
||||
pred = preds[0].squeeze()
|
||||
pred = (pred > sensitivity).float()
|
||||
|
||||
mask = transforms.ToPILImage()(pred)
|
||||
mask = mask.resize(orig_image.size)
|
||||
|
||||
if mask_blur > 0:
|
||||
mask = mask.filter(ImageFilter.GaussianBlur(radius=mask_blur))
|
||||
|
||||
if mask_offset != 0:
|
||||
from PIL import ImageMorph
|
||||
if mask_offset > 0:
|
||||
pattern = [[1,1,1],[1,1,1],[1,1,1]]
|
||||
for _ in range(mask_offset):
|
||||
mask = mask.filter(ImageFilter.MaxFilter(3))
|
||||
else:
|
||||
for _ in range(-mask_offset):
|
||||
mask = mask.filter(ImageFilter.MinFilter(3))
|
||||
|
||||
new_im = orig_image.copy()
|
||||
new_im.putalpha(mask)
|
||||
|
||||
new_im_tensor = pil2tensor(new_im)
|
||||
mask_tensor = pil2tensor(mask)
|
||||
|
||||
processed_images.append(new_im_tensor)
|
||||
processed_masks.append(mask_tensor)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
new_ims = torch.cat(processed_images, dim=0)
|
||||
new_masks = torch.cat(processed_masks, dim=0)
|
||||
|
||||
return (new_ims, new_masks)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AILAB_RMBG": AILAB_RMBG
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AILAB_RMBG": "🧽 RMBG (Remove Background)"
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=61.0"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "ComfyUI-RMBG"
|
||||
version = "1.0.0"
|
||||
description = "A ComfyUI node for removing image backgrounds using RMBG-2.0"
|
||||
authors = [{ name = "1038 Lab" }]
|
||||
license = { text = "MIT" }
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"torch>=2.0.0,<3.0.0",
|
||||
"torchvision>=0.15.0,<1.0.0",
|
||||
"Pillow>=9.0.0,<10.0.0",
|
||||
"numpy>=1.22.0,<2.0.0",
|
||||
"transformers>=4.30.0,<5.0.0",
|
||||
"safetensors>=0.3.0,<1.0.0"
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
include = ["ComfyUI-RMBG*"]
|
||||
@@ -0,0 +1,6 @@
|
||||
torch>=2.0.0,<3.0.0
|
||||
torchvision>=0.15.0,<1.0.0
|
||||
Pillow>=9.0.0,<10.0.0
|
||||
numpy>=1.22.0,<2.0.0
|
||||
transformers>=4.30.0,<5.0.0
|
||||
safetensors>=0.3.0,<1.0.0
|
||||
Reference in New Issue
Block a user