Only download the nsfw detector when used
This commit is contained in:
@@ -72,19 +72,6 @@ if not os.path.exists(REACTOR_MODELS_PATH):
|
||||
if not os.path.exists(FACE_MODELS_PATH):
|
||||
os.makedirs(FACE_MODELS_PATH)
|
||||
|
||||
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/confusion_matrix.png",
|
||||
"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)
|
||||
download(model_url, model_path, model_name)
|
||||
|
||||
dir_facerestore_models = os.path.join(models_dir, "facerestore_models")
|
||||
os.makedirs(dir_facerestore_models, exist_ok=True)
|
||||
folder_paths.folder_names_and_paths["facerestore_models"] = ([dir_facerestore_models], folder_paths.supported_pt_extensions)
|
||||
@@ -196,7 +183,7 @@ class reactor:
|
||||
global FACE_SIZE, FACE_HELPER
|
||||
|
||||
self.face_helper = FACE_HELPER
|
||||
|
||||
|
||||
faceSize = 512
|
||||
if "1024" in face_restore_model.lower():
|
||||
faceSize = 1024
|
||||
@@ -233,7 +220,7 @@ class reactor:
|
||||
sd = comfy.utils.load_torch_file(model_path, safe_load=True)
|
||||
facerestore_model = model_loading.load_state_dict(sd).eval()
|
||||
facerestore_model.to(device)
|
||||
|
||||
|
||||
if faceSize != FACE_SIZE or self.face_helper is None:
|
||||
self.face_helper = FaceRestoreHelper(1, face_size=faceSize, crop_ratio=(1, 1), det_model=facedetection, save_ext='png', use_parse=True, device=device)
|
||||
FACE_SIZE = faceSize
|
||||
@@ -330,7 +317,7 @@ class reactor:
|
||||
result = restored_img_tensor
|
||||
|
||||
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):
|
||||
|
||||
if face_boost is not None:
|
||||
@@ -356,7 +343,7 @@ class reactor:
|
||||
|
||||
if face_model == "none":
|
||||
face_model = None
|
||||
|
||||
|
||||
script = FaceSwapScript()
|
||||
pil_images = batch_tensor_to_pil(input_image)
|
||||
|
||||
@@ -413,7 +400,7 @@ class reactor:
|
||||
|
||||
if self.restore or not self.face_boost_enabled:
|
||||
result = reactor.restore_face(self,result,face_restore_model,face_restore_visibility,codeformer_weight,facedetection)
|
||||
|
||||
|
||||
else:
|
||||
image_black = Image.new("RGB", (512, 512))
|
||||
result = batched_pil_to_tensor([image_black])
|
||||
@@ -428,7 +415,7 @@ class ReActorPlusOpt:
|
||||
return {
|
||||
"required": {
|
||||
"enabled": ("BOOLEAN", {"default": True, "label_off": "OFF", "label_on": "ON"}),
|
||||
"input_image": ("IMAGE",),
|
||||
"input_image": ("IMAGE",),
|
||||
"swap_model": (list(model_names().keys()),),
|
||||
"facedetection": (["retinaface_resnet50", "retinaface_mobile0.25", "YOLOv5l", "YOLOv5n"],),
|
||||
"face_restore_model": (get_model_names(get_restorers),),
|
||||
@@ -462,7 +449,7 @@ class ReActorPlusOpt:
|
||||
self.interpolation = "Bicubic"
|
||||
self.boost_model_visibility = 1
|
||||
self.boost_cf_weight = 0.5
|
||||
|
||||
|
||||
def execute(self, enabled, input_image, swap_model, facedetection, face_restore_model, face_restore_visibility, codeformer_weight, source_image=None, face_model=None, options=None, face_boost=None):
|
||||
|
||||
if options is not None:
|
||||
@@ -472,13 +459,13 @@ class ReActorPlusOpt:
|
||||
self.detect_gender_source = options["detect_gender_source"]
|
||||
self.input_faces_index = options["input_faces_index"]
|
||||
self.source_faces_index = options["source_faces_index"]
|
||||
|
||||
|
||||
if face_boost is not None:
|
||||
self.face_boost_enabled = face_boost["enabled"]
|
||||
self.restore = face_boost["restore_with_main_after"]
|
||||
else:
|
||||
self.face_boost_enabled = False
|
||||
|
||||
|
||||
result = reactor.execute(
|
||||
self,enabled,input_image,swap_model,self.detect_gender_source,self.detect_gender_input,self.source_faces_index,self.input_faces_index,self.console_log_level,face_restore_model,face_restore_visibility,codeformer_weight,facedetection,source_image,face_model,self.faces_order, face_boost=face_boost
|
||||
)
|
||||
@@ -494,7 +481,7 @@ class LoadFaceModel:
|
||||
"face_model": (get_model_names(get_facemodels),),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("FACE_MODEL",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "🌌 ReActor"
|
||||
@@ -513,7 +500,7 @@ class LoadFaceModel:
|
||||
class BuildFaceModel:
|
||||
def __init__(self):
|
||||
self.output_dir = FACE_MODELS_PATH
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -551,14 +538,14 @@ class BuildFaceModel:
|
||||
face_model = analyze_faces(image, det_size_half)
|
||||
if face_model is not None and len(face_model) > 0:
|
||||
print("...........................................................", end=" ")
|
||||
|
||||
|
||||
if face_model is not None and len(face_model) > 0:
|
||||
return face_model[0]
|
||||
else:
|
||||
no_face_msg = "No face found, please try another image"
|
||||
# logger.error(no_face_msg)
|
||||
return no_face_msg
|
||||
|
||||
|
||||
def blend_faces(self, save_mode, send_only, face_model_name, compute_method, images=None, face_models=None):
|
||||
global BLENDED_FACE_MODEL
|
||||
blended_face: Face = BLENDED_FACE_MODEL
|
||||
@@ -589,7 +576,7 @@ class BuildFaceModel:
|
||||
print(f"{int(((i+1)/n)*100)}%")
|
||||
faces.append(face)
|
||||
embeddings.append(face.embedding)
|
||||
|
||||
|
||||
elif face_models is not None:
|
||||
|
||||
n = len(face_models)
|
||||
@@ -695,7 +682,7 @@ class RestoreFace:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"image": ("IMAGE",),
|
||||
"facedetection": (["retinaface_resnet50", "retinaface_mobile0.25", "YOLOv5l", "YOLOv5n"],),
|
||||
"model": (get_model_names(get_restorers),),
|
||||
"visibility": ("FLOAT", {"default": 1, "min": 0.0, "max": 1, "step": 0.05}),
|
||||
@@ -734,7 +721,7 @@ class MaskHelper:
|
||||
# self.force_resize_width = 0
|
||||
# self.force_resize_height = 0
|
||||
# self.resize_behavior = "source_size"
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
bboxs = ["bbox/"+x for x in folder_paths.get_filename_list("ultralytics_bbox")]
|
||||
@@ -764,7 +751,7 @@ class MaskHelper:
|
||||
"mask_optional": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE","MASK","IMAGE","IMAGE")
|
||||
RETURN_NAMES = ("IMAGE","MASK","MASK_PREVIEW","SWAPPED_FACE")
|
||||
FUNCTION = "execute"
|
||||
@@ -792,7 +779,7 @@ class MaskHelper:
|
||||
if len(self.labels) > 0:
|
||||
segs, _ = masking_segs.filter(segs, self.labels)
|
||||
# segs, _ = masking_segs.filter(segs, "all")
|
||||
|
||||
|
||||
sam_modelname = folder_paths.get_full_path("sams", sam_model_name)
|
||||
|
||||
if 'vit_h' in sam_model_name:
|
||||
@@ -813,12 +800,12 @@ class MaskHelper:
|
||||
sam.is_auto_mode = self.device_mode == "AUTO"
|
||||
|
||||
combined_mask, _ = core.make_sam_mask_segmented(sam, segs, images, self.detection_hint, sam_dilation, sam_threshold, bbox_expansion, mask_hint_threshold, mask_hint_use_negative)
|
||||
|
||||
|
||||
else:
|
||||
combined_mask = mask_optional
|
||||
|
||||
# *** MASK TO IMAGE ***:
|
||||
|
||||
|
||||
mask_image = combined_mask.reshape((-1, 1, combined_mask.shape[-2], combined_mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3)
|
||||
|
||||
# *** MASK MORPH ***:
|
||||
@@ -835,12 +822,12 @@ class MaskHelper:
|
||||
elif morphology_operation == "close":
|
||||
mask_image = self.dilate(mask_image, morphology_distance)
|
||||
mask_image = self.erode(mask_image, morphology_distance)
|
||||
|
||||
|
||||
# *** MASK BLUR ***:
|
||||
|
||||
|
||||
if len(mask_image.size()) == 3:
|
||||
mask_image = mask_image.unsqueeze(3)
|
||||
|
||||
|
||||
mask_image = mask_image.permute(0, 3, 1, 2)
|
||||
kernel_size = blur_radius * 2 + 1
|
||||
sigma = sigma_factor * (0.6 * blur_radius - 0.3)
|
||||
@@ -849,7 +836,7 @@ class MaskHelper:
|
||||
mask_image_final = mask_image_final[:, :, :, 0]
|
||||
|
||||
# *** CUT BY MASK ***:
|
||||
|
||||
|
||||
if len(swapped_image.shape) < 4:
|
||||
C = 1
|
||||
else:
|
||||
@@ -906,7 +893,7 @@ class MaskHelper:
|
||||
single = (swapped_image[i, ymin:ymax+1, xmin:xmax+1,:]).unsqueeze(0)
|
||||
resized = torch.nn.functional.interpolate(single.permute(0, 3, 1, 2), size=(use_height, use_width), mode='bicubic').permute(0, 2, 3, 1)
|
||||
cutted_image[i] = resized[0]
|
||||
|
||||
|
||||
# Preserve our type unless we were previously RGB and added non-opaque alpha due to the mask size
|
||||
if C == 1:
|
||||
cutted_image = core.tensor2mask(cutted_image)
|
||||
@@ -959,7 +946,7 @@ class MaskHelper:
|
||||
|
||||
result = image_base.detach().clone()
|
||||
face_segment = mask_image_final
|
||||
|
||||
|
||||
for i in range(0, MB):
|
||||
if is_empty[i]:
|
||||
continue
|
||||
@@ -1038,7 +1025,7 @@ class MaskHelper:
|
||||
face_segment[...,3] = mask[i]
|
||||
|
||||
result = rgba2rgb_tensor(result)
|
||||
|
||||
|
||||
return (result,combined_mask,mask_image_final,face_segment,)
|
||||
|
||||
def gaussian_blur(self, image, kernel_size, sigma):
|
||||
@@ -1067,7 +1054,7 @@ class MaskHelper:
|
||||
output_tensor = output_reshaped.reshape(batch_size, num_channels, height, width)
|
||||
|
||||
return output_tensor
|
||||
|
||||
|
||||
def erode(self, image, distance):
|
||||
return 1. - self.dilate(1. - image, distance)
|
||||
|
||||
@@ -1084,7 +1071,7 @@ class ImageDublicator:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"image": ("IMAGE",),
|
||||
"count": ("INT", {"default": 1, "min": 0}),
|
||||
},
|
||||
}
|
||||
@@ -1096,7 +1083,7 @@ class ImageDublicator:
|
||||
CATEGORY = "🌌 ReActor"
|
||||
|
||||
def execute(self, image, count):
|
||||
images = [image for i in range(count)]
|
||||
images = [image for i in range(count)]
|
||||
return (images,)
|
||||
|
||||
|
||||
@@ -1114,7 +1101,7 @@ class ImageRGBA2RGB:
|
||||
CATEGORY = "🌌 ReActor"
|
||||
|
||||
def execute(self, image):
|
||||
out = rgba2rgb_tensor(image)
|
||||
out = rgba2rgb_tensor(image)
|
||||
return (out,)
|
||||
|
||||
|
||||
@@ -1123,7 +1110,7 @@ class MakeFaceModelBatch:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"face_model1": ("FACE_MODEL",),
|
||||
"face_model1": ("FACE_MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"face_model2": ("FACE_MODEL",),
|
||||
@@ -1217,7 +1204,7 @@ class ReActorFaceBoost:
|
||||
"restore_with_main_after": restore_with_main_after,
|
||||
}
|
||||
return (face_boost, )
|
||||
|
||||
|
||||
class ReActorUnload:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
@@ -1,12 +1,30 @@
|
||||
from transformers import pipeline
|
||||
from PIL import Image
|
||||
import logging
|
||||
import os
|
||||
from reactor_utils import download
|
||||
|
||||
def ensure_nsfw_model(model_path):
|
||||
"""Download NSFW detection model if it doesn't exist"""
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(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/confusion_matrix.png",
|
||||
"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(model_path, model_name)
|
||||
download(model_url, model_path, model_name)
|
||||
|
||||
SCORE = 0.965 # 0.965 and less - is safety content
|
||||
|
||||
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)
|
||||
result = predict(img)
|
||||
|
||||
Reference in New Issue
Block a user