refactor: Update WatermarkDetectionNode to use torchvision EfficientNet - Replace efficientnet_pytorch with torchvision implementation
This commit is contained in:
@@ -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",
|
||||
}
|
||||
Reference in New Issue
Block a user