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:
+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