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:
Gourieff | 古仁
2025-05-21 12:47:35 +07:00
parent 0addca8a40
commit f6b4a0ebce
4 changed files with 97 additions and 27 deletions
+30 -11
View File
@@ -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"
+9
View File
@@ -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
View File
@@ -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
+21 -1
View File
@@ -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: