UPDs and FIXs
- #25 Fix - ComfyUI native ProgressBar for different steps - ORIGINAL_IMAGE output for main nodes - no tmp file for nsfw detector - nsfw detector little speed up
This commit is contained in:
@@ -12,6 +12,7 @@ import cv2
|
||||
import math
|
||||
from typing import List
|
||||
from PIL import Image
|
||||
import io
|
||||
from scipy import stats
|
||||
from insightface.app.common import Face
|
||||
from segment_anything import sam_model_registry
|
||||
@@ -50,7 +51,9 @@ from reactor_utils import (
|
||||
prepare_cropped_face,
|
||||
normalize_cropped_face,
|
||||
add_folder_path_and_extensions,
|
||||
rgba2rgb_tensor
|
||||
rgba2rgb_tensor,
|
||||
progress_bar,
|
||||
progress_bar_reset
|
||||
)
|
||||
from reactor_patcher import apply_patch
|
||||
from r_facelib.utils.face_restoration_helper import FaceRestoreHelper
|
||||
@@ -153,7 +156,8 @@ class reactor:
|
||||
"hidden": {"faces_order": "FACES_ORDER"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","FACE_MODEL")
|
||||
RETURN_TYPES = ("IMAGE","FACE_MODEL","IMAGE")
|
||||
RETURN_NAMES = ("SWAPPED_IMAGE","FACE_MODEL","ORIGINAL_IMAGE")
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "🌌 ReActor"
|
||||
|
||||
@@ -233,10 +237,12 @@ class reactor:
|
||||
|
||||
out_images = []
|
||||
|
||||
pbar = progress_bar(total_images)
|
||||
|
||||
for i in range(total_images):
|
||||
|
||||
if total_images > 1:
|
||||
logger.status(f"Restoring {i+1}")
|
||||
logger.status(f"Restoring {i}")
|
||||
|
||||
cur_image_np = image_np[i,:, :, ::-1]
|
||||
|
||||
@@ -311,16 +317,25 @@ class reactor:
|
||||
if state.interrupted or model_management.processing_interrupted():
|
||||
logger.status("Interrupted by User")
|
||||
return input_image
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
restored_img_np = np.array(out_images).astype(np.float32) / 255.0
|
||||
restored_img_tensor = torch.from_numpy(restored_img_np)
|
||||
|
||||
result = restored_img_tensor
|
||||
|
||||
progress_bar_reset(pbar)
|
||||
|
||||
return result
|
||||
|
||||
def execute(self, enabled, input_image, swap_model, detect_gender_source, detect_gender_input, source_faces_index, input_faces_index, console_log_level, face_restore_model,face_restore_visibility, codeformer_weight, facedetection, source_image=None, face_model=None, faces_order=None, face_boost=None):
|
||||
|
||||
device = model_management.get_torch_device()
|
||||
|
||||
if isinstance(input_image, torch.Tensor) and input_image.device != device:
|
||||
input_image = input_image.to(device)
|
||||
|
||||
if face_boost is not None:
|
||||
self.face_boost_enabled = face_boost["enabled"]
|
||||
self.boost_model = face_boost["boost_model"]
|
||||
@@ -349,20 +364,21 @@ class reactor:
|
||||
pil_images = batch_tensor_to_pil(input_image)
|
||||
|
||||
# NSFW checker
|
||||
logger.status("Checking for any unsafe content")
|
||||
pbar = progress_bar(len(pil_images))
|
||||
pil_images_sfw = []
|
||||
tmp_img = "reactor_tmp.png"
|
||||
for img in pil_images:
|
||||
if state.interrupted or model_management.processing_interrupted():
|
||||
logger.status("Interrupted by User")
|
||||
break
|
||||
img.save(tmp_img)
|
||||
if not sfw.nsfw_image(tmp_img, NSFWDET_MODEL_PATH):
|
||||
img_byte_arr = io.BytesIO()
|
||||
img.save(img_byte_arr, format='PNG')
|
||||
img_byte_arr = img_byte_arr.getvalue()
|
||||
if not sfw.nsfw_image(img_byte_arr, NSFWDET_MODEL_PATH):
|
||||
pil_images_sfw.append(img)
|
||||
if os.path.exists(tmp_img):
|
||||
os.remove(tmp_img)
|
||||
pbar.update(1)
|
||||
pil_images = pil_images_sfw
|
||||
# # #
|
||||
progress_bar_reset(pbar)
|
||||
|
||||
if len(pil_images) > 0:
|
||||
|
||||
@@ -392,6 +408,7 @@ class reactor:
|
||||
interpolation=self.interpolation,
|
||||
)
|
||||
result = batched_pil_to_tensor(p.init_images)
|
||||
original_image = batched_pil_to_tensor(pil_images)
|
||||
|
||||
if face_model is None:
|
||||
current_face_model = get_current_faces_model()
|
||||
@@ -406,8 +423,9 @@ class reactor:
|
||||
image_black = Image.new("RGB", (512, 512))
|
||||
result = batched_pil_to_tensor([image_black])
|
||||
face_model_to_provide = None
|
||||
original_image = result
|
||||
|
||||
return (result,face_model_to_provide)
|
||||
return (result,face_model_to_provide,original_image)
|
||||
|
||||
|
||||
class ReActorPlusOpt:
|
||||
@@ -431,7 +449,8 @@ class ReActorPlusOpt:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE","FACE_MODEL")
|
||||
RETURN_TYPES = ("IMAGE","FACE_MODEL","IMAGE")
|
||||
RETURN_NAMES = ("SWAPPED_IMAGE","FACE_MODEL","ORIGINAL_IMAGE")
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "🌌 ReActor"
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ import urllib.request
|
||||
import onnxruntime
|
||||
from typing import Any
|
||||
import folder_paths
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
ORT_SESSION = None
|
||||
|
||||
@@ -236,6 +237,14 @@ def normalize_cropped_face(cropped_face):
|
||||
return cropped_face
|
||||
|
||||
|
||||
def progress_bar(total):
|
||||
return ProgressBar(total)
|
||||
|
||||
def progress_bar_reset(pbar):
|
||||
pbar.current = 0
|
||||
pbar.update(0)
|
||||
|
||||
|
||||
# author: Trung0246 --->
|
||||
def add_folder_path_and_extensions(folder_name, full_folder_paths, extensions):
|
||||
# Iterate over the list of full folder paths
|
||||
|
||||
+37
-15
@@ -1,34 +1,56 @@
|
||||
from transformers import pipeline
|
||||
from PIL import Image
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
import comfy.model_management as model_management
|
||||
from reactor_utils import download
|
||||
from scripts.reactor_logger import logger
|
||||
|
||||
MODEL_EXISTS = False
|
||||
|
||||
def ensure_nsfw_model(nsfwdet_model_path):
|
||||
"""Download NSFW detection model if it doesn't exist"""
|
||||
if not os.path.exists(nsfwdet_model_path):
|
||||
os.makedirs(nsfwdet_model_path)
|
||||
nd_urls = [
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/config.json",
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/model.safetensors",
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/preprocessor_config.json",
|
||||
]
|
||||
for model_url in nd_urls:
|
||||
model_name = os.path.basename(model_url)
|
||||
model_path = os.path.join(nsfwdet_model_path, model_name)
|
||||
global MODEL_EXISTS
|
||||
downloaded = 0
|
||||
nd_urls = [
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/config.json",
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/model.safetensors",
|
||||
"https://huggingface.co/AdamCodd/vit-base-nsfw-detector/resolve/main/preprocessor_config.json",
|
||||
]
|
||||
for model_url in nd_urls:
|
||||
model_name = os.path.basename(model_url)
|
||||
model_path = os.path.join(nsfwdet_model_path, model_name)
|
||||
if not os.path.exists(model_path):
|
||||
if not os.path.exists(nsfwdet_model_path):
|
||||
os.makedirs(nsfwdet_model_path)
|
||||
download(model_url, model_path, model_name)
|
||||
if os.path.exists(model_path):
|
||||
downloaded += 1
|
||||
MODEL_EXISTS = True if downloaded == 3 else False
|
||||
return MODEL_EXISTS
|
||||
|
||||
SCORE = 0.96
|
||||
|
||||
logging.getLogger("transformers").setLevel(logging.ERROR)
|
||||
|
||||
def nsfw_image(img_path: str, model_path: str):
|
||||
ensure_nsfw_model(model_path)
|
||||
with Image.open(img_path) as img:
|
||||
predict = pipeline("image-classification", model=model_path)
|
||||
def nsfw_image(img_data, model_path: str):
|
||||
if not MODEL_EXISTS:
|
||||
logger.status("Ensuring NSFW detection model exists...")
|
||||
if not ensure_nsfw_model(model_path):
|
||||
return True
|
||||
logger.status("Checking for any unsafe content...")
|
||||
device = model_management.get_torch_device()
|
||||
with Image.open(io.BytesIO(img_data)) as img:
|
||||
if "cpu" in str(device):
|
||||
predict = pipeline("image-classification", model=model_path)
|
||||
else:
|
||||
device_id = 0
|
||||
if "cuda" in str(device):
|
||||
device_id = int(str(device).split(":")[1])
|
||||
predict = pipeline("image-classification", model=model_path, device=device_id)
|
||||
result = predict(img)
|
||||
if result[0]["label"] == "nsfw" and result[0]["score"] > SCORE:
|
||||
logger.status(f"NSFW content detected, skipping...")
|
||||
logger.status(f"NSFW content detected with score={result[0]["score"]}, skipping...")
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -22,6 +22,8 @@ from scripts.reactor_logger import logger
|
||||
from reactor_utils import (
|
||||
move_path,
|
||||
get_image_md5hash,
|
||||
progress_bar,
|
||||
progress_bar_reset
|
||||
)
|
||||
from scripts.r_faceboost import swapper, restorer
|
||||
|
||||
@@ -179,7 +181,12 @@ def half_det_size(det_size):
|
||||
|
||||
def analyze_faces(img_data: np.ndarray, det_size=(640, 640)):
|
||||
face_analyser = getAnalysisModel(det_size)
|
||||
faces = face_analyser.get(img_data)
|
||||
|
||||
faces = []
|
||||
try:
|
||||
faces = face_analyser.get(img_data)
|
||||
except:
|
||||
logger.error("No faces found")
|
||||
|
||||
# Try halving det_size if no faces are found
|
||||
if len(faces) == 0 and det_size[0] > 320 and det_size[1] > 320:
|
||||
@@ -463,6 +470,8 @@ def swap_face_many(
|
||||
if source_faces is not None:
|
||||
|
||||
target_faces = []
|
||||
pbar = progress_bar(len(target_imgs))
|
||||
|
||||
for i, target_img in enumerate(target_imgs):
|
||||
if state.interrupted or model_management.processing_interrupted():
|
||||
logger.status("Interrupted by User")
|
||||
@@ -504,7 +513,11 @@ def swap_face_many(
|
||||
# target_face = analyze_faces(target_img)
|
||||
if target_face is not None:
|
||||
target_faces.append(target_face)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
progress_bar_reset(pbar)
|
||||
|
||||
# No use in trying to swap faces if no faces are found, enhancement
|
||||
if len(target_faces) == 0:
|
||||
logger.status("Cannot detect any Target, skipping swapping...")
|
||||
@@ -527,6 +540,8 @@ def swap_face_many(
|
||||
|
||||
source_face_idx = 0
|
||||
|
||||
pbar = progress_bar(len(target_imgs))
|
||||
|
||||
for face_num in faces_index:
|
||||
# No use in trying to swap faces if no further faces are found, enhancement
|
||||
if face_num >= len(target_faces):
|
||||
@@ -554,12 +569,15 @@ def swap_face_many(
|
||||
# logger.status(f"Swapping as-is")
|
||||
result = face_swapper.get(target_img, target_face_single, source_face)
|
||||
results[i] = result
|
||||
pbar.update(1)
|
||||
elif wrong_gender == 1:
|
||||
wrong_gender = 0
|
||||
logger.status("Wrong target gender detected")
|
||||
pbar.update(1)
|
||||
continue
|
||||
else:
|
||||
logger.status(f"No target face found for {face_num}")
|
||||
pbar.update(1)
|
||||
elif src_wrong_gender == 1:
|
||||
src_wrong_gender = 0
|
||||
logger.status("Wrong source gender detected")
|
||||
@@ -567,6 +585,8 @@ def swap_face_many(
|
||||
else:
|
||||
logger.status(f"No source face found for face number {source_face_idx}.")
|
||||
|
||||
progress_bar_reset(pbar)
|
||||
|
||||
result_images = [Image.fromarray(cv2.cvtColor(result, cv2.COLOR_BGR2RGB)) for result in results]
|
||||
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user