Files
gorillaframeai-GF_nodes/gfrbmg2plus.py
T
2025-04-19 17:40:06 +03:00

322 lines
12 KiB
Python

import torch
import torch.nn.functional as F
import folder_paths
import os
from PIL import Image
from transformers import AutoModelForImageSegmentation
from torchvision.transforms.functional import normalize
import numpy as np
import cv2
from huggingface_hub import hf_hub_download, HfApi
device = "cuda" if torch.cuda.is_available() else "cpu"
# Добавляем путь к модели RMBG-2.0
folder_paths.add_model_folder_path("rmbg_models", os.path.join(folder_paths.models_dir, "RMBG", "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)
def resize_image(image, interpolation_method):
image = image.convert('RGB')
w, h = image.size
new_w = (w + 63) // 64 * 64
new_h = (h + 63) // 64 * 64
image = image.resize((new_w, new_h), interpolation_method)
return image
def center_crop_or_pad(image, target_size, scale=1.0, rotation=0.0):
actual_scale = 1.0 + (scale / 10.0)
target_w, target_h = target_size
w, h = image.size
width_ratio = target_w / w
new_w = target_w
new_h = int(h * width_ratio)
new_w = int(new_w * actual_scale)
new_h = int(new_h * actual_scale)
image = image.resize((new_w, new_h), Image.LANCZOS)
if rotation != 0:
image = image.rotate(-rotation, expand=True)
w, h = image.size
left = (w - target_w) // 2
top = (h - target_h) // 2
right = left + target_w
bottom = top + target_h
if left < 0 or top < 0 or right > w or bottom > h:
new_image = Image.new(image.mode, target_size, (0, 0, 0, 0))
paste_left = max(0, -left)
paste_top = max(0, -top)
crop_left = max(0, left)
crop_top = max(0, top)
crop_right = min(w, right)
crop_bottom = min(h, bottom)
new_image.paste(image.crop((crop_left, crop_top, crop_right, crop_bottom)),
(paste_left, paste_top))
return new_image
else:
return image.crop((left, top, right, bottom))
class GFrbmg2Plus:
_model_instance = None
@classmethod
def get_model(cls):
if cls._model_instance is None:
cls.initialize_model()
return cls._model_instance
@classmethod
def initialize_model(cls):
model_path = os.path.join(folder_paths.models_dir, "RMBG", "RMBG-2.0")
if not os.path.exists(model_path):
os.makedirs(model_path, exist_ok=True)
api = HfApi()
files = api.list_repo_files("1038lab/RMBG-2.0")
for file in files:
if '/' not in file:
if file.endswith(('.json', '.py', '.safetensors')):
print(f"Downloading {file}...")
hf_hub_download(
repo_id="1038lab/RMBG-2.0",
filename=file,
local_dir=model_path
)
cls._model_instance = AutoModelForImageSegmentation.from_pretrained(
model_path,
trust_remote_code=True,
local_files_only=True
)
cls._model_instance.to(device)
cls._model_instance.eval()
print("Model loaded successfully.")
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"invert_mask": ("BOOLEAN", {"default": False}),
"expand_mask": ("FLOAT", {"default": 0.0, "min": -255, "max": 255, "step": 0.1, "display": "number"}),
"blur_weight": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 255, "step": 0.1, "display": "number"}),
"sticker_size": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 255, "step": 0.1, "display": "number"}),
"sticker_blur": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 255, "step": 0.1, "display": "number"}),
"sticker_color": ("COLOR", {"default": "#000000"}),
"background_color": ("COLOR", {"default": "#000000"}),
"use_original_bg": ("BOOLEAN", {"default": False}),
"rotation": ("FLOAT", {
"default": 0.0,
"min": -360.0,
"max": 360.0,
"step": 0.1,
"display": "number"
}),
# Теперь этот параметр строго в конце required-блока
"interpolation_method": (["Lanczos", "Bicubic", "Bilinear", "Nearest"], {"default": "Lanczos"})
},
"optional": {
"bg_image": ("IMAGE",),
"bg_image_scale": ("FLOAT", {
"default": 0.0,
"min": -10.0,
"max": 10.0,
"step": 0.1,
"display": "number"
})
}
}
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "MASK")
RETURN_NAMES = ("image_rgba", "mask", "image", "sticker_mask")
FUNCTION = "remove_background"
CATEGORY = "🐵 GorillaFrame/Image"
def __init__(self):
self.model = self.get_model()
def clean_mask(self, mask, expand_mask, blur_weight):
try:
if expand_mask == 0 and blur_weight == 0:
return mask
mask_np = np.array(mask)
mask_tensor = torch.from_numpy(mask_np).float().cuda()
if mask_tensor is None or mask_tensor.numel() == 0:
return mask
if blur_weight > 0:
print("Applying blur to mask...")
kernel_size = max(3, int(blur_weight * 20) | 1)
sigma = blur_weight * 5
padding = kernel_size // 2
mask_blur = mask_tensor.unsqueeze(0).unsqueeze(0)
coords = torch.arange(kernel_size, device='cuda').float() - padding
x, y = torch.meshgrid(coords, coords)
kernel_2d = torch.exp(-(x.pow(2) + y.pow(2)) / (2 * sigma ** 2))
kernel_2d = kernel_2d / kernel_2d.sum()
kernel_2d = kernel_2d.view(1, 1, kernel_size, kernel_size)
mask_blur = F.conv2d(mask_blur, kernel_2d, padding=padding)
mask_tensor = mask_blur.squeeze()
if expand_mask != 0:
print("Expanding mask...")
kernel_size = max(3, int(abs(expand_mask) * 10) | 1)
padding = kernel_size // 2
if expand_mask > 0:
pool = torch.nn.MaxPool2d(kernel_size, stride=1, padding=padding)
mask_tensor = pool(mask_tensor.unsqueeze(0)).squeeze(0)
else:
pool = torch.nn.MaxPool2d(kernel_size, stride=1, padding=padding)
mask_tensor = -pool(-mask_tensor.unsqueeze(0)).squeeze(0)
mask_np = mask_tensor.cpu().numpy().astype(np.uint8)
del mask_tensor
torch.cuda.empty_cache()
return Image.fromarray(mask_np)
except Exception as e:
print(f"Error in clean_mask: {e}")
return mask
def remove_background(self, image, invert_mask, expand_mask, blur_weight, sticker_size, sticker_blur, sticker_color, background_color, bg_image=None, bg_image_scale=1.0, use_original_bg=False, rotation=0, interpolation_method="Lanczos"):
print("Starting background removal process...")
def parse_color(color):
if isinstance(color, bool):
return (0, 0, 0)
color_str = str(color).lstrip('#')
return tuple(int(color_str[i:i+2], 16) for i in (0, 2, 4))
sticker_rgb = parse_color(sticker_color)
background_rgb = parse_color(background_color)
processed_images = []
processed_masks = []
processed_blacks = []
processed_sticker_masks = []
center_pil = None
if bg_image is not None and len(bg_image) > 0:
center_pil = tensor2pil(bg_image[0])
for img in image:
orig_image = tensor2pil(img)
w, h = orig_image.size
if w < 64 or h < 64:
print(f"Warning: Image size {w}x{h} is too small, might cause issues")
max_size = 4096
if w > max_size or h > max_size:
print(f"Warning: Image size {w}x{h} is very large, might cause memory issues")
rgba_image = Image.new('RGBA', (w, h), (0, 0, 0, 0))
if use_original_bg:
black_image = orig_image.copy()
else:
if bg_image is not None and len(bg_image) > 0:
bg_pil = tensor2pil(bg_image[0])
black_image = center_crop_or_pad(bg_pil, (w, h), bg_image_scale, rotation)
else:
black_image = Image.new('RGB', (w, h), background_rgb)
black_image.paste(rgba_image, (0, 0), rgba_image)
interpolation_mapping = {
"Lanczos": Image.LANCZOS,
"Bicubic": Image.BICUBIC,
"Bilinear": Image.BILINEAR,
"Nearest": Image.NEAREST
}
interpolation = interpolation_mapping.get(interpolation_method, Image.LANCZOS)
image_resized = resize_image(orig_image, interpolation)
im_np = np.array(image_resized)
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
im_tensor = torch.unsqueeze(im_tensor,0)
im_tensor = torch.divide(im_tensor,255.0)
im_tensor = normalize(im_tensor,[0.485, 0.456, 0.406],[0.229, 0.224, 0.225])
if torch.cuda.is_available():
im_tensor=im_tensor.cuda()
with torch.no_grad():
print("Running model inference...")
result = self.model(im_tensor)[-1].sigmoid().cpu()
result = result[0].squeeze()
result = F.interpolate(result.unsqueeze(0).unsqueeze(0), size=(h, w), mode='bilinear').squeeze()
mask_pil = tensor2pil(result)
mask_pil = self.clean_mask(mask_pil, expand_mask, blur_weight)
if invert_mask:
print("Inverting mask...")
mask_np = np.array(mask_pil)
mask_np = 255 - mask_np
mask_pil = Image.fromarray(mask_np)
if orig_image.mode != 'RGBA':
orig_image = orig_image.convert('RGBA')
rgba_image.paste(orig_image, (0, 0), mask_pil)
black_image.paste(orig_image, (0, 0), mask_pil)
if sticker_size > 0:
print("Processing sticker...")
mask_np = np.array(mask_pil)
_, binary_mask = cv2.threshold(mask_np, 127, 255, cv2.THRESH_BINARY)
edges = cv2.Canny(binary_mask, 100, 200)
contours, _ = cv2.findContours(edges, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)
main_contour_mask = np.zeros_like(mask_np)
cv2.drawContours(main_contour_mask, contours, -1, 255, int(sticker_size), lineType=cv2.LINE_AA)
if sticker_blur > 0:
blur_size = int(sticker_blur * 2) * 2 + 1
main_contour_mask = cv2.GaussianBlur(main_contour_mask, (blur_size, blur_size), 0)
sticker_mask = Image.fromarray(main_contour_mask)
sticker = Image.new('RGBA', orig_image.size, (*sticker_rgb, 255))
rgba_image.paste(sticker, mask=sticker_mask)
black_sticker = Image.new('RGB', orig_image.size, sticker_rgb)
black_image.paste(black_sticker, mask=sticker_mask)
processed_sticker_masks.append(pil2tensor(sticker_mask))
else:
empty_mask = Image.new('L', orig_image.size, 0)
processed_sticker_masks.append(pil2tensor(empty_mask))
processed_images.append(pil2tensor(rgba_image))
processed_masks.append(pil2tensor(mask_pil))
processed_blacks.append(pil2tensor(black_image))
print("Background removal process completed.")
return (torch.cat(processed_images, dim=0),
torch.cat(processed_masks, dim=0),
torch.cat(processed_blacks, dim=0),
torch.cat(processed_sticker_masks, dim=0))
NODE_CLASS_MAPPINGS = {
"GFrbmg2Plus": GFrbmg2Plus
}
NODE_DISPLAY_NAME_MAPPINGS = {
"GFrbmg2Plus": "🐵 GF Remove Background (GFrbmg2Plus)"
}