From 4a7587a3feabc835ef5d7daf5b1a77acba29d654 Mon Sep 17 00:00:00 2001 From: Paul Date: Mon, 27 Jan 2025 15:46:52 +0000 Subject: [PATCH] Only download the nsfw detector when used --- nodes.py | 79 ++++++++++++++++++------------------------ scripts/reactor_sfw.py | 18 ++++++++++ 2 files changed, 51 insertions(+), 46 deletions(-) diff --git a/nodes.py b/nodes.py index b4e65bf..c9745c6 100644 --- a/nodes.py +++ b/nodes.py @@ -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): diff --git a/scripts/reactor_sfw.py b/scripts/reactor_sfw.py index 524b074..f91cc66 100644 --- a/scripts/reactor_sfw.py +++ b/scripts/reactor_sfw.py @@ -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)