diff --git a/nodes/ImageCaptioningNode.py b/nodes/ImageCaptioningNode.py index cfe8660..c91da07 100644 --- a/nodes/ImageCaptioningNode.py +++ b/nodes/ImageCaptioningNode.py @@ -4,6 +4,12 @@ from transformers import BlipProcessor,AutoModel, BlipForConditionalGeneration,A from PIL import Image import numpy as np from scipy.ndimage import binary_dilation +import torchvision.transforms as transforms +import os +import requests +from pathlib import Path +import folder_paths +import torchvision.models as models class ImageCaptioningNode: @@ -152,10 +158,10 @@ class ArtOrHumanClassifierNode: proba = outputs.logits.softmax(1) # Get the probabilities for "artificial" and "human" classes - artificial_prob = proba[0][0].item() - human_prob = proba[0][1].item() + artificial_prob = float(proba[0][0].item()) + human_prob = float(proba[0][1].item()) - output_ui = {"text": [artificial_prob]} if show_on_node else {} + output_ui = {"text": [f"Artificial: {artificial_prob:.2%}\nHuman: {human_prob:.2%}"]} if show_on_node else {} return {"result": (artificial_prob, human_prob), "ui": output_ui} @@ -167,7 +173,7 @@ class DocumentClassificationNode: def INPUT_TYPES(cls): return {"required": {"image": ("IMAGE", {"default": None})}} - RETURN_TYPES = ("INT", "STRING") + RETURN_TYPES = ("FLOAT", "STRING") FUNCTION = "classify" CATEGORY = "LexTools/ImageProcessing/Classification" @@ -186,12 +192,230 @@ class DocumentClassificationNode: # Perform the classification outputs = self.model(**inputs) logits = outputs.logits + probabilities = torch.softmax(logits, dim=1) predicted_class_index = torch.argmax(logits, dim=1).item() + confidence_score = float(probabilities[0][predicted_class_index].item()) # Get the class name predicted_class_name = self.class_names[predicted_class_index] - return (predicted_class_index, predicted_class_name) + return (confidence_score, predicted_class_name) + +class NSFWClassifierNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", {"default": None}), + "show_on_node": ("BOOLEAN", {"default": False}), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}), + }, + } + OUTPUT_NODE = True + + RETURN_TYPES = ("STRING", "FLOAT", "FLOAT", "BOOLEAN", "BOOLEAN") # Added boolean outputs + RETURN_NAMES = ("Classification", "SFW Score", "NSFW Score", "Is SFW", "Is NSFW") + FUNCTION = "classify_nsfw" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def __init__(self): + self.feature_extractor = AutoFeatureExtractor.from_pretrained("umairrkhn/fine-tuned-nsfw-classification") + self.model = AutoModelForImageClassification.from_pretrained("umairrkhn/fine-tuned-nsfw-classification") + + def classify_nsfw(self, image, show_on_node, threshold): + try: + # Convert the image tensor to numpy array + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + + # Process the image + inputs = self.feature_extractor(images=img, return_tensors="pt") + outputs = self.model(**inputs) + probs = outputs.logits.softmax(1)[0] + + # Get probabilities for each class (model has 2 classes: SFW and NSFW) + sfw_prob = float(probs[0].item()) # SFW + nsfw_prob = float(probs[1].item()) # NSFW + + # Get the predicted class + predicted_class_idx = probs.argmax().item() + class_names = ["SFW", "NSFW"] + predicted_class = class_names[predicted_class_idx] + + # Determine boolean states using threshold + is_sfw = sfw_prob >= threshold + is_nsfw = nsfw_prob >= threshold + + # Format the results string + results = f"Predicted: {predicted_class}\n" + results += f"SFW: {sfw_prob:.2%} ({'Yes' if is_sfw else 'No'})\n" + results += f"NSFW: {nsfw_prob:.2%} ({'Yes' if is_nsfw else 'No'})" + + output_ui = {"text": [results]} if show_on_node else {} + + return {"result": (results, sfw_prob, nsfw_prob, is_sfw, is_nsfw), + "ui": output_ui} + + except Exception as e: + print(f"Error in NSFW classification: {str(e)}") + return {"result": (str(e), 0.0, 0.0, False, False), + "ui": {"text": [str(e)]} if show_on_node else {}} + +class WatermarkDetectionNode: + model = None # Class-level model instance for caching + transform = None # Class-level transform for caching + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", {"default": None}), + "show_on_node": ("BOOLEAN", {"default": False}), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}), + }, + } + OUTPUT_NODE = True + + RETURN_TYPES = ("STRING", "FLOAT", "FLOAT", "BOOLEAN", "BOOLEAN") + RETURN_NAMES = ("Classification", "Clean Score", "Watermark Score", "Is Clean", "Has Watermark") + FUNCTION = "detect_watermark" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def download_model(self): + # Create models directory if it doesn't exist + models_dir = os.path.join(os.path.dirname(os.path.dirname(__file__)), "models") + os.makedirs(models_dir, exist_ok=True) + + model_path = os.path.join(models_dir, "watermark_model.pt") + + # Download the model if it doesn't exist + if not os.path.exists(model_path): + print("Downloading watermark detection model...") + url = "https://huggingface.co/qwertyforce/watermark_detection/resolve/main/model.pt" + try: + response = requests.get(url, stream=True) + response.raise_for_status() + + with open(model_path, 'wb') as f: + for chunk in response.iter_content(chunk_size=8192): + f.write(chunk) + print("Model downloaded successfully") + except Exception as e: + print(f"Error downloading model: {str(e)}") + # Try alternative URL from scenery_watermarks repo + url = "https://huggingface.co/qwertyforce/scenery_watermarks/resolve/main/model.pt" + print("Trying alternative model source...") + response = requests.get(url, stream=True) + response.raise_for_status() + + with open(model_path, 'wb') as f: + for chunk in response.iter_content(chunk_size=8192): + f.write(chunk) + print("Model downloaded successfully from alternative source") + + return model_path + + def __init__(self): + if WatermarkDetectionNode.transform is None: + # Standard EfficientNet preprocessing + WatermarkDetectionNode.transform = transforms.Compose([ + transforms.Resize(256), + transforms.CenterCrop(224), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) + ]) + + if WatermarkDetectionNode.model is None: + try: + # Create a new EfficientNet model + base_model = models.efficientnet_b0(pretrained=False) + # Modify the classifier for 2 classes + base_model.classifier = torch.nn.Sequential( + torch.nn.Dropout(p=0.2, inplace=True), + torch.nn.Linear(in_features=1280, out_features=2, bias=True) + ) + + # Download and load the state dict + model_path = self.download_model() + state_dict = torch.load(model_path, map_location='cpu') + + # If it's a state dict, try to load it + if isinstance(state_dict, dict): + try: + # Try direct loading + base_model.load_state_dict(state_dict) + except: + try: + # Try removing 'module.' prefix + new_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()} + base_model.load_state_dict(new_state_dict) + except Exception as e: + print(f"Failed to load state dict: {str(e)}") + # If both attempts fail, just use the base model + pass + else: + # If it's already a model, try to extract its state dict + try: + base_model.load_state_dict(state_dict.state_dict()) + except: + print("Failed to load model state dict, using base model") + + WatermarkDetectionNode.model = base_model + if torch.cuda.is_available(): + WatermarkDetectionNode.model = WatermarkDetectionNode.model.cuda() + WatermarkDetectionNode.model.eval() + + except Exception as e: + print(f"Error loading watermark detection model: {str(e)}") + raise + + def detect_watermark(self, image, show_on_node, threshold): + try: + # Convert the image tensor to PIL Image + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + + # Ensure image is RGB + if img.mode != 'RGB': + img = img.convert('RGB') + + # Preprocess the image + img_tensor = self.transform(img).unsqueeze(0) + if torch.cuda.is_available(): + img_tensor = img_tensor.cuda() + + # Get model predictions + with torch.no_grad(): + outputs = self.model(img_tensor) + probs = torch.softmax(outputs, dim=1)[0] + + # Get probabilities for each class + clean_prob = float(probs[0].item()) # Clean image + watermark_prob = float(probs[1].item()) # Watermarked image + + # Get the predicted class + predicted_class_idx = probs.argmax().item() + class_names = ["Clean", "Watermarked"] + predicted_class = class_names[predicted_class_idx] + + # Determine boolean states using threshold + is_clean = clean_prob >= threshold + has_watermark = watermark_prob >= threshold + + # Format the results string + results = f"Predicted: {predicted_class}\n" + results += f"Clean: {clean_prob:.2%} ({'Yes' if is_clean else 'No'})\n" + results += f"Watermarked: {watermark_prob:.2%} ({'Yes' if has_watermark else 'No'})" + + output_ui = {"text": [results]} if show_on_node else {} + + return {"result": (results, clean_prob, watermark_prob, is_clean, has_watermark), + "ui": output_ui} + + except Exception as e: + print(f"Error in watermark detection: {str(e)}") + return {"result": (str(e), 0.0, 0.0, False, False), + "ui": {"text": [str(e)]} if show_on_node else {}} NODE_CLASS_MAPPINGS = { "AgeClassifierNode": AgeClassifierNode, @@ -199,10 +423,14 @@ NODE_CLASS_MAPPINGS = { "DocumentClassificationNode": DocumentClassificationNode, "ImageCaptioning": ImageCaptioningNode, "ArtOrHumanClassifierNode": ArtOrHumanClassifierNode, + "NSFWClassifierNode": NSFWClassifierNode, + "WatermarkDetectionNode": WatermarkDetectionNode, } NODE_DISPLAY_NAME_MAPPINGS = { "ImageScaleToMin": "Image Scale To Min", "ImageCaptioning": "Image Captioning", "ArtOrHumanClassifierNode": "Art Or Human Classifier", + "NSFWClassifierNode": "NSFW Classifier", + "WatermarkDetectionNode": "Watermark Detector", } \ No newline at end of file