diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..5fc0224 --- /dev/null +++ b/__init__.py @@ -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" +] \ No newline at end of file diff --git a/ailab_RMBG.py b/ailab_RMBG.py new file mode 100644 index 0000000..e76cc85 --- /dev/null +++ b/ailab_RMBG.py @@ -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)" +} \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..e59c6c0 --- /dev/null +++ b/pyproject.toml @@ -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*"] \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..e6b0910 --- /dev/null +++ b/requirements.txt @@ -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 \ No newline at end of file