Author SHA1 Message Date
Craig WrightandCopilot 8fb9b27ea1 Update nodes/SegformerNode.py
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2025-03-26 21:44:41 +00:00
Craig Wright 877dd93b0c fix: Resolve tensor dimension mismatch in SegformerNode - Improve mask handling and preview generation - Add better error handling and dimension checks - Fix image and mask tensor format compatibility 2025-03-26 21:41:34 +00:00
Craig WrightandCopilot 6f3b85c157 Update README.md
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2025-03-26 17:45:57 +00:00
Craig WrightandCopilot 12cff98987 Update README.md
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2025-03-26 17:45:49 +00:00
Craig Wright 9d64aad7cc docs: Update README.md to reflect changes in NSFWClassifierNode and WatermarkDetectionNode inputs and outputs, and add torchvision to installation instructions 2025-03-26 17:34:38 +00:00
Craig Wright 024fdf44b4 chore: Bump version to 1.0.2 in pyproject.toml 2025-03-26 17:26:15 +00:00
Craig Wright 78a0fe1ddb feat: Update WatermarkDetectionNode with improved model handling and error recovery 2025-03-25 03:28:12 +00:00
Craig Wright e1b49cc738 refactor: Update WatermarkDetectionNode to use torchvision EfficientNet - Replace efficientnet_pytorch with torchvision implementation 2025-03-25 03:26:49 +00:00
Craig Wright d7ff05266f Enhance SegformerNodeMergeSegments with new mask processing options and preview functionality 2025-03-25 02:39:09 +00:00
Craig Wright 7ae11ac705 Add tracking file for project structure and key files 2025-03-25 00:40:29 +00:00
Craig Wright d1da468459 Update pyproject.toml 2025-03-24 22:56:11 +00:00
Craig Wright f17a163bff Merge pull request #12 from 6DEADSHOT9/main
Fast Api issue fixed
2025-03-24 22:53:26 +00:00
6DEADSHOT9 75bf61118b Fast Api issue fixed 2024-09-24 19:55:43 +05:30
Craig Wright ae3b49a80e Update pyproject.toml 2024-06-28 20:25:13 +01:00
Craig Wright 2f75f817e4 Update requirements.txt 2024-06-28 20:03:04 +01:00
Craig Wright 4e4a19185b Merge branch 'main' into pr/7 2024-06-28 20:00:01 +01:00
Craig Wright 9dbd068a71 Merge pull request #8 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-06-28 19:53:44 +01:00
Craig Wright 2dfc5cfbe5 Merge pull request #9 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-05-23 10:39:07 +01:00
haohaocreates 438df02026 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-22 17:13:41 -04:00
haohaocreates 285ed93e76 chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-22 17:13:38 -04:00
Srmsamay 3341bcea6e Corrected FastApi import statement 2024-03-15 23:13:14 +05:30
Craig Wright a559d3815d Merge pull request #4 from nidefawl/main
Fix boolean types and class mappings
2024-02-19 18:53:39 +00:00
nidefawl 533910af92 Fix boolean types 2023-12-17 16:48:50 +01:00
nidefawl 55dccd5943 Add missing class mappings 2023-12-17 16:48:33 +01:00
Craig Wright 3cff522e9a Merge pull request #3 from alpertunga-bile/main
automatic installation for packages
2023-11-18 04:41:03 +00:00
alpertunga-bile a41973d0c3 automatic installation update is added 2023-10-23 22:41:38 +03:00
11 changed files with 1087 additions and 662 deletions
+21
View File
@@ -0,0 +1,21 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+10
View File
@@ -0,0 +1,10 @@
.github/workflows/publish.yml
.gitignore
README.md
__init__.py
nodes/ImageCaptioningNode.py
nodes/ImageProcessingNode.py
nodes/SegformerNode.py
nodes/__init__.py
pyproject.toml
requirements.txt
+77 -42
View File
@@ -25,65 +25,100 @@ ComfyUI-LexTools is a Python-based image processing and analysis toolkit that us
- _Output_: Converted score. - _Output_: Converted score.
Additional nodes from [GitHub Pages](https://github.com/strimmlarn/ComfyUI-Strimmlarns-Aesthetic-Score/) - These have been modified to improve performance and add an option to store the model in RAM, which significantly reduces generation time: Additional nodes from [GitHub Pages](https://github.com/strimmlarn/ComfyUI-Strimmlarns-Aesthetic-Score/) - These have been modified to improve performance and add an option to store the model in RAM, which significantly reduces generation time:
- `CalculateAestheticScore`: An optimized version of the original, with an option to keep the model loaded in RAM. (No specific input or output detailed in the provided code) - `CalculateAestheticScore`: An optimized version of the original, with an option to keep the model loaded in RAM.
- `AesthetlcScoreSorter`: Sorts the images by score. (No specific input or output detailed in the provided code) - `AestheticScoreSorter`: Sorts the images by score.
- `AesteticModel`: Loads the aesthetic model. (No specific input or output detailed in the provided code) - `AestheticModel`: Loads the aesthetic model.
2. **ImageCaptioningNode.py** - Implements nodes for image captioning and classification: 2. **ImageCaptioningNode.py** - Implements nodes for image captioning and classification:
- `ImageCaptioningNode`: Provides a caption for the image. - `ImageCaptioningNode`: Provides a caption for the image using BLIP model.
- _Input_: `image` (IMAGE) - _Input_: `image` (IMAGE)
- _Output_: String caption. - _Output_: String caption.
- `FoodCategoryNode`: Classifies the food category of an image. - `FoodCategoryClassifierNode`: Classifies food categories in images.
- _Input_: `image` (IMAGE) - _Input_: `image` (IMAGE)
- _Output_: String category. - _Output_: Top 5 food categories with probabilities.
- `AgeClassifierNode`: Classifies the age of a person in the image. - `AgeClassifierNode`: Classifies the age range in images.
- _Input_: `image` (IMAGE) - _Input_: `image` (IMAGE)
- _Output_: String age range. - _Output_: Top 5 age ranges with probabilities.
- `ImageClassifierNode`: General image classification. - `ArtOrHumanClassifierNode`: Detects if an image is AI-generated or human-made.
- _Input_: `image` (IMAGE), `show_on_node` (BOOL) - _Input_: `image` (IMAGE), `show_on_node` (BOOL)
- _Output_: String label, `artificial_prob` (INT), `human_prob` (INT) - _Output_: Artificial and human probabilities.
- `ClassifierNode`: A generic classifier node. - `DocumentClassificationNode`: Classifies document types.
- _Input_: `image` (IMAGE) - _Input_: `image` (IMAGE)
- _Output_: String label. - _Output_: Document type index and name.
- `NSFWClassifierNode`: Classifies content safety levels.
- _Input_: `image` (IMAGE), `show_on_node` (BOOL), `threshold` (FLOAT)
- _Output_:
- Classification report (STRING)
- SFW Score (FLOAT)
- NSFW Score (FLOAT)
- Is SFW (BOOLEAN)
- Is NSFW (BOOLEAN)
- `WatermarkDetectionNode`: Detects watermarks in images using EfficientNet.
- _Input_: `image` (IMAGE), `show_on_node` (BOOL), `threshold` (FLOAT)
- _Output_:
- Classification report (STRING)
- Clean Score (FLOAT)
- Watermark Score (FLOAT)
- Is Clean (BOOLEAN)
- Has Watermark (BOOLEAN)
3. **SegformerNode.py** - Handles semantic segmentation of images. It includes various nodes such as: 3. **SegformerNode.py** - Handles semantic segmentation of images:
- `SegformerNode`: Performs segmentation of the image. - `SegformerNode`: Performs semantic segmentation with multiple model options.
- _Input_: `image` (IMAGE), `model_name` (STRING), `show_on_node` (BOOL) - _Input_: `image` (IMAGE), `model_name` (STRING), `normalize_mask` (BOOL), `binary_mask` (BOOL), `resize_mode` (STRING), `invert_mask` (BOOL), `show_preview` (BOOL), `return_individual_masks` (BOOL), `post_process` (STRING), `post_process_radius` (INT), `segment_groups` (STRING)
- _Output_: Segmented image. - _Output_: Segmented image, mask, info, and preview.
- `SegformerNodeMasks`: Provides masks for the segmented images. - `SegformerNodeMasks`: Creates individual segment masks.
- _Input_: No specific input detailed in the provided code. - _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING)
- _Output_: Image masks. - _Output_: Image, mask, and segment info.
- `SegformerNodeMergeSegments`: Merges certain segments in the segmented image. - `SegformerNodeMergeSegments`: Merges and processes segments with advanced options.
- _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING), `blur_radius` (INT), `dilation_radius` (INT), `intensity` (INT), `ceiling` (INT), `show_on_node` (BOOL) - _Input_: `image` (IMAGE), `segments_to_merge_str` (STRING), `model_name` (STRING), `normalize_mask` (BOOL), `binary_mask` (BOOL), `resize_mode` (STRING), `invert_mask` (BOOL), `show_preview` (BOOL), `blur_radius` (INT), `dilation_radius` (INT), `intensity` (FLOAT), `ceiling` (FLOAT)
- _Output_: Image with merged segments. - _Output_: Processed image, mask, info, and preview.
- `SeedIncrementerNode`: Increment the seed used for random processes. - `SeedIncrementerNode`: Manages seed incrementation for workflows.
- _Input_: `seed` (INT), `increment_at` (INT) - _Input_: `seed` (INT), `IncrementAt` (INT)
- _Output_: Incremented seed. - _Output_: Seed string, seed int, subseed string, subseed int.
- `StepCfgIncrementNode`: Calculates the step configuration for the process. - `StepCfgIncrementNode`: Handles step and configuration increments.
- _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `img_steps` (INT), `max_steps` (INT) - _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `image_steps` (INT), `max_steps` (INT)
- _Output_: Calculated step configuration. - _Output_: CFG and steps values.
## Requirements ## Requirements
The project primarily uses the following libraries: The project requires the following Python libraries:
- Python - torch
- Torch - transformers
- Transformers - Pillow (PIL)
- PIL - matplotlib
- Matplotlib - numpy
- Numpy - scipy
- IO - huggingface_hub
- Scipy - torchvision
## Installation ## Installation
To install the necessary libraries, run: 1. Install the required Python packages:
```bash ```bash
pip install torch transformers pillow matplotlib numpy scipy pip install torch transformers pillow matplotlib numpy scipy huggingface_hub torchvision
``` ```
2. Clone this repository into your ComfyUI custom_nodes directory:
```bash
cd ComfyUI/custom_nodes
git clone https://github.com/YourUsername/ComfyUI-LexTools.git
```
3. Restart ComfyUI to load the new nodes.
## Usage
The nodes will appear in the ComfyUI interface under the "LexTools" category, organized into subcategories:
- LexTools/ImageProcessing/Segmentation
- LexTools/ImageProcessing/Classification
- LexTools/ImageProcessing/Captioning
- LexTools/Utilities
## Contributing ## Contributing
Contributions to this project are welcome. If you find a bug or think of a feature that would benefit the project, please open an issue. If you'd like to contribute code, please open a pull request.
Contributions are welcome! Please feel free to submit a Pull Request. For major changes, please open an issue first to discuss what you would like to change.
## License
This project is licensed under the MIT License - see the LICENSE file for details.
+1
View File
@@ -1,5 +1,6 @@
#from .nodes.SegGPT import segGPTNode #from .nodes.SegGPT import segGPTNode
from .nodes import SegformerNode,ImageCaptioningNode,ImageProcessingNode from .nodes import SegformerNode,ImageCaptioningNode,ImageProcessingNode
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
**SegformerNode.NODE_CLASS_MAPPINGS, **SegformerNode.NODE_CLASS_MAPPINGS,
Binary file not shown.
+234 -6
View File
@@ -4,6 +4,12 @@ from transformers import BlipProcessor,AutoModel, BlipForConditionalGeneration,A
from PIL import Image from PIL import Image
import numpy as np import numpy as np
from scipy.ndimage import binary_dilation 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: class ImageCaptioningNode:
@@ -127,7 +133,7 @@ class ArtOrHumanClassifierNode:
return { return {
"required": { "required": {
"image": ("IMAGE",), "image": ("IMAGE",),
"show_on_node": ("BOOL", {"default": False}), "show_on_node": ("BOOLEAN", {"default": False}),
}, },
} }
OUTPUT_NODE = True OUTPUT_NODE = True
@@ -152,10 +158,10 @@ class ArtOrHumanClassifierNode:
proba = outputs.logits.softmax(1) proba = outputs.logits.softmax(1)
# Get the probabilities for "artificial" and "human" classes # Get the probabilities for "artificial" and "human" classes
artificial_prob = proba[0][0].item() artificial_prob = float(proba[0][0].item())
human_prob = proba[0][1].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} return {"result": (artificial_prob, human_prob), "ui": output_ui}
@@ -167,7 +173,7 @@ class DocumentClassificationNode:
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return {"required": {"image": ("IMAGE", {"default": None})}} return {"required": {"image": ("IMAGE", {"default": None})}}
RETURN_TYPES = ("INT", "STRING") RETURN_TYPES = ("FLOAT", "STRING")
FUNCTION = "classify" FUNCTION = "classify"
CATEGORY = "LexTools/ImageProcessing/Classification" CATEGORY = "LexTools/ImageProcessing/Classification"
@@ -186,12 +192,230 @@ class DocumentClassificationNode:
# Perform the classification # Perform the classification
outputs = self.model(**inputs) outputs = self.model(**inputs)
logits = outputs.logits logits = outputs.logits
probabilities = torch.softmax(logits, dim=1)
predicted_class_index = torch.argmax(logits, dim=1).item() predicted_class_index = torch.argmax(logits, dim=1).item()
confidence_score = float(probabilities[0][predicted_class_index].item())
# Get the class name # Get the class name
predicted_class_name = self.class_names[predicted_class_index] 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 = { NODE_CLASS_MAPPINGS = {
"AgeClassifierNode": AgeClassifierNode, "AgeClassifierNode": AgeClassifierNode,
@@ -199,10 +423,14 @@ NODE_CLASS_MAPPINGS = {
"DocumentClassificationNode": DocumentClassificationNode, "DocumentClassificationNode": DocumentClassificationNode,
"ImageCaptioning": ImageCaptioningNode, "ImageCaptioning": ImageCaptioningNode,
"ArtOrHumanClassifierNode": ArtOrHumanClassifierNode, "ArtOrHumanClassifierNode": ArtOrHumanClassifierNode,
"NSFWClassifierNode": NSFWClassifierNode,
"WatermarkDetectionNode": WatermarkDetectionNode,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"ImageScaleToMin": "Image Scale To Min", "ImageScaleToMin": "Image Scale To Min",
"ImageCaptioning": "Image Captioning", "ImageCaptioning": "Image Captioning",
"ArtOrHumanClassifierNode": "Art Or Human Classifier", "ArtOrHumanClassifierNode": "Art Or Human Classifier",
"NSFWClassifierNode": "NSFW Classifier",
"WatermarkDetectionNode": "Watermark Detector",
} }
+237 -478
View File
@@ -1,19 +1,12 @@
import hashlib import hashlib
import fastapi from fastapi import FastAPI
import fastapi import torch, time
import torch
import time
import io import io
import cv2
from transformers import AutoModelForCausalLM, AutoTokenizer
from torchvision import transforms
import comfy.samplers import comfy.samplers
import matplotlib.transforms as mpl_transforms from matplotlib import transforms
from PIL import Image, ImageFilter, ImageEnhance, ImageOps, ImageDraw, ImageChops, ImageFont from PIL import Image, ImageFilter, ImageEnhance, ImageOps, ImageDraw, ImageChops, ImageFont
import numpy as np import numpy as np
from scipy.ndimage import zoom
import comfy.model_management as model_management import comfy.model_management as model_management
import json import json
import uuid import uuid
@@ -24,10 +17,8 @@ import torch.nn as nn
from os.path import join from os.path import join
import clip import clip
import folder_paths import folder_paths
# create path to aesthetic model. # create path to aesthetic model.
folder_paths.folder_names_and_paths["aesthetic"] = ([os.path.join( folder_paths.folder_names_and_paths["aesthetic"] = ([os.path.join(folder_paths.models_dir,"aesthetic")], folder_paths.supported_pt_extensions)
folder_paths.models_dir, "aesthetic")], folder_paths.supported_pt_extensions)
aspect_ratios = [ aspect_ratios = [
@@ -41,34 +32,24 @@ aspect_ratios = [
MAX_RESOLUTION = 10240 # adjust this value as needed MAX_RESOLUTION = 10240 # adjust this value as needed
# Tensor to PIL # Tensor to PIL
def tensor2pil(image): def tensor2pil(image):
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
# PIL to Tensor # PIL to Tensor
def pil2tensor(image): def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
# PIL Hex # PIL Hex
def pil2hex(image): def pil2hex(image):
return hashlib.sha256(np.array(tensor2pil(image)).astype(np.uint16).tobytes()).hexdigest() return hashlib.sha256(np.array(tensor2pil(image)).astype(np.uint16).tobytes()).hexdigest()
# PIL to Mask # PIL to Mask
def pil2mask(image): def pil2mask(image):
image_np = np.array(image.convert("L")).astype(np.float32) / 255.0 image_np = np.array(image.convert("L")).astype(np.float32) / 255.0
mask = torch.from_numpy(image_np) mask = torch.from_numpy(image_np)
return 1.0 - mask return 1.0 - mask
# Mask to PIL # Mask to PIL
def mask2pil(mask): def mask2pil(mask):
if mask.ndim > 2: if mask.ndim > 2:
mask = mask.squeeze(0) mask = mask.squeeze(0)
@@ -77,199 +58,6 @@ def mask2pil(mask):
return mask_pil return mask_pil
def scale_and_print_mask(mask, target_shape=(11, 11)):
"""
Scale a given mask to a target shape and print it rounded to 4 decimal places.
"""
# Calculate scaling factors
scale_x = target_shape[0] / mask.shape[0]
scale_y = target_shape[1] / mask.shape[1]
# Rescale the mask
scaled_mask = zoom(mask, (scale_x, scale_y))
# Print the scaled mask, rounded to 4 decimal places
for row in scaled_mask:
print(", ".join([f"{x:.2f}" for x in row]))
def apply_feathering(mask, feathering_distance=10):
# Apply Gaussian blur to the mask
mask_feathered = mask
if feathering_distance > 0:
# Generate kernel
kernel_size = 2 * feathering_distance + 1
kernel = cv2.getGaussianKernel(kernel_size, feathering_distance)
# Convert the mask tensor to a numpy array
mask_np = mask.numpy()
# Apply Gaussian blur to the mask
kernel_2d = np.dot(kernel, kernel.T)
mask_feathered_np = cv2.filter2D(
mask_np, -1, kernel_2d, borderType=cv2.BORDER_CONSTANT)
# Convert the result back to a PyTorch tensor
mask_feathered = torch.tensor(mask_feathered_np, dtype=torch.float32)
return mask_feathered
def apply_gradient(mask, transition_points, feathering_distance):
# Ensure feathering_distance is at least 1
gradient = torch.linspace(0, 1, max(feathering_distance, 1))
for x, y in transition_points:
if x + feathering_distance < mask.shape[0]:
mask[x: x + feathering_distance, y] = gradient
if x - feathering_distance >= 0:
mask[x - feathering_distance: x, y] = gradient[::-1]
return mask
class ImageAspectPadNode:
@classmethod
def INPUT_TYPES(cls):
global aspect_ratios # Assuming aspect_ratios is a list of aspect ratio strings
return {
"required": {
"image": ("IMAGE",),
"aspect_ratio": (aspect_ratios, {"default": aspect_ratios[0]}),
"invert_ratio": (["true", "false"], {"default": "false"}),
"edge_feathering_distance": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"feathering": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"exclude_out_of_bounds": (["true", "false"], {"default": "false"}),
"left_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"right_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"top_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"bottom_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
},
"optional": {
"show_on_node": ("INT", {"default": 0}),
}
}
RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "expand_image"
OUTPUT_NODE = True
CATEGORY = "LexTools/ImageProcessing/AspectPad"
def expand_image(self, image, aspect_ratio, invert_ratio, feathering, left_padding, right_padding, top_padding, bottom_padding, show_on_node, exclude_out_of_bounds, edge_feathering_distance):
debug_info = {} # Initialize debug info dictionary
debug = True
try:
# Initial setup
d1, d2, d3, d4 = image.size()
aspect_ratio = float(aspect_ratio.split(
'/')[0]) / float(aspect_ratio.split('/')[1])
if invert_ratio == "true":
aspect_ratio = 1.0 / aspect_ratio
# Padding calculations
image_aspect_ratio = d3 / d2
if image_aspect_ratio > aspect_ratio:
pad_height = int(d3 / aspect_ratio) - d2
top_padding += pad_height // 2
bottom_padding += pad_height - top_padding
else:
pad_width = int(d2 * aspect_ratio) - d3
left_padding += pad_width // 2
right_padding += pad_width - left_padding
# Debug Information
debug_info['image_size'] = (d1, d2, d3, d4)
debug_info['padding'] = (
top_padding, bottom_padding, left_padding, right_padding)
# Identify the mask boundary
boundary_top = top_padding
boundary_bottom = top_padding + d2
boundary_left = left_padding
boundary_right = left_padding + d3
# Initialize new image and mask
new_image = torch.zeros((d1, d2 + top_padding + bottom_padding,
d3 + left_padding + right_padding, d4), dtype=torch.float32)
new_image[:, top_padding:top_padding + d2,
left_padding:left_padding + d3, :] = image
mask = torch.ones((d2 + top_padding + bottom_padding,
d3 + left_padding + right_padding), dtype=torch.float32)
mask[top_padding:top_padding + d2,
left_padding:left_padding + d3] = 0
if debug == True:
scale_and_print_mask(mask)
# Apply edge feathering to the identified boundary within mask
transition_points = []
for i in range(1, mask.shape[0]):
for j in range(mask.shape[1]):
if mask[i, j] != mask[i-1, j]:
transition_points.append((i, j))
for i in range(mask.shape[0]):
for j in range(1, mask.shape[1]):
if mask[i, j] != mask[i, j-1]:
transition_points.append((i, j))
if debug == True:
print("Transition Points:", transition_points) # Debug line
# Check if "exclude_out_of_bounds" is set to "true"
if exclude_out_of_bounds == "true":
# Create a mask that excludes the areas touching the bounds
inner_mask = torch.zeros_like(mask)
inner_mask[boundary_top:boundary_bottom, boundary_left:boundary_right] = 1
inner_mask = 1 - inner_mask # Invert the inner mask
# Apply the inner mask to the original mask
mask = mask * inner_mask
# Filter transition points to only include those within the bounds
transition_points = [point for point in transition_points if boundary_top <= point[0] < boundary_bottom and boundary_left <= point[1] < boundary_right]
# Apply edge feathering to the identified boundary within the new mask
if edge_feathering_distance > 0:
try:
mask = apply_gradient(mask, transition_points, edge_feathering_distance)
except Exception as e:
if debug == True:
print("An error occurred:", str(e))
# Apply edge feathering to the identified boundary within mask
if debug == True:
print("Mask After Before Feather:") # Debug line
# Assuming this function prints the mask
scale_and_print_mask(mask)
# Apply overall feathering
if feathering > 0:
try:
mask = apply_feathering(mask, feathering)
except Exception as e:
# scale_and_print_mask(mask)
if debug == True:
print("An error occurred:", str(e))
if debug == True:
print("Mask After Feather:") # Debug line
# Assuming this function prints the mask
scale_and_print_mask(mask)
# Debugging output
print("Debug Information:", debug_info)
output_ui = {}
if show_on_node == 1:
output_ui = {"ui": {"images": [new_image]}}
return (new_image, mask, output_ui)
except Exception as e:
print("An error occurred:", str(e))
if debug == True:
print("Debug Information:", debug_info)
raise # Re-raise the caught exception for further handling
class ImageRankingNode: class ImageRankingNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -314,63 +102,117 @@ class ImageRankingNode:
json.dump(data, f) json.dump(data, f)
class AutoModelForCausalLMNode: class ImageAspectPadNode:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { global aspect_ratios # Assuming aspect_ratios is a list of aspect ratio strings
"MESSAGE": ("STRING",), return {
"MaxTokens": ("INTEGER",) "required": {
}} "image": ("IMAGE",),
"aspect_ratio": (aspect_ratios, {"default": aspect_ratios[0]}),
"invert_ratio": (["true", "false"], {"default": "false"}),
"feathering": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"left_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"right_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"top_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"bottom_padding": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
RETURN_TYPES = ("STRING") },
FUNCTION = "caption" "optional": {
"show_on_node": ("INT", {"default": 0}),
}
}
CATEGORY = "LexTools/TextGeneration" RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "expand_image"
OUTPUT_NODE = True
def __init__(self): CATEGORY = "LexTools/ImageProcessing/AspectPad"
self.tokenizer = AutoTokenizer.from_pretrained(
"mistralai/Mistral-7B-Instruct-v0.1")
self.model = AutoModelForCausalLM.from_pretrained(
"mistralai/Mistral-7B-Instruct-v0.1")
def caption(self, MESSAGE, MaxTokens): def expand_image(self, image, aspect_ratio, invert_ratio, feathering, left_padding, right_padding, top_padding, bottom_padding,show_on_node):
d1, d2, d3, d4 = image.size()
aspect_ratio = float(aspect_ratio.split('/')[0]) / float(aspect_ratio.split('/')[1])
if invert_ratio == "true":
aspect_ratio = 1.0 / aspect_ratio
device = "cuda" # the device to load the model onto image_aspect_ratio = d3 / d2
messages = [ if image_aspect_ratio > aspect_ratio:
{"role": "user", "content": MESSAGE}, pad_height = int(d3 / aspect_ratio) - d2
top_padding += pad_height // 2
bottom_padding += pad_height - top_padding
else:
pad_width = int(d2 * aspect_ratio) - d3
left_padding += pad_width // 2
right_padding += pad_width - left_padding
new_image = torch.zeros(
(d1, d2 + top_padding + bottom_padding, d3 + left_padding + right_padding, d4),
dtype=torch.float32,
)
new_image[:, top_padding:top_padding + d2, left_padding:left_padding + d3, :] = image
] mask = torch.ones(
encodeds = self.tokenizer.apply_chat_template( (d2 + top_padding + bottom_padding, d3 + left_padding + right_padding),
messages, return_tensors="pt") dtype=torch.float32,
)
model_inputs = encodeds.to(device) t = torch.zeros(
self.model.to(device) (d2, d3),
dtype=torch.float32
)
generated_ids = self.model.generate( if feathering > 0 and feathering * 2 < d2 and feathering * 2 < d3:
model_inputs, max_new_tokens=MaxTokens, do_sample=True)
decoded = self.tokenizer.batch_decode(generated_ids)
return (decoded[0]) for i in range(d2):
for j in range(d3):
dt = i if top_padding != 0 else d2
db = d2 - i if bottom_padding != 0 else d2
dl = j if left_padding != 0 else d3
dr = d3 - j if right_padding != 0 else d3
d = min(dt, db, dl, dr)
if d >= feathering:
continue
v = (feathering - d) / feathering
t[i, j] = v * v
mask[top_padding:top_padding + d2, left_padding:left_padding + d3] = t
output_ui = {}
if show_on_node ==1:
output_ui = {"ui": {"images": [new_image]}}
return (new_image, mask, output_ui)
class ImageScaleToMin: class ImageScaleToMin:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": {"image": ("IMAGE",)}, return {"required": {"image": ("IMAGE",)},
"optional": {"MinScalePix": ("FLOAT", {"default": 512, "min": 0.0, "max": 2056, "step": 1}), }} "optional":{"MinScalePix": ("FLOAT", {"default": 512, "min": 0.0, "max": 2056, "step": 1}),}}
RETURN_TYPES = ("FLOAT",) RETURN_TYPES = ("FLOAT",)
FUNCTION = "calculate_scale" FUNCTION = "calculate_scale"
CATEGORY = "LexTools/ImageProcessing/upscaling" CATEGORY = "LexTools/ImageProcessing/upscaling"
def calculate_scale(self, image, MinScalePix): def calculate_scale(self, image,MinScalePix):
d1, height, width, d4 = image.shape d1, height, width, d4 = image.shape
min_dim = min(width, height) min_dim = min(width, height)
scale = MinScalePix / min_dim scale = MinScalePix / min_dim
return (scale,) return (scale,)
class ImageFilterByIntScoreNode: class ImageFilterByIntScoreNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -391,70 +233,91 @@ class ImageFilterByIntScoreNode:
if score < threshold: if score < threshold:
pass pass
else: else:
return (image,) return (image,)
class ImageFilterByFloatScoreNode: class ImageFilterByFloatScoreNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"score": ("FLOAT", {"default": 0.0}), "image": ("IMAGE",),
"threshold": ("FLOAT", {"default": 0.0}), "score": ("FLOAT", {"default": 0.0, "min": -100.0, "max": 100.0}),
"image": ("IMAGE", {"default": None}), "threshold": ("FLOAT", {"default": 5.0, "min": -100.0, "max": 100.0}),
"show_on_node": ("BOOLEAN", {"default": False}),
}, },
} }
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ("IMAGE", "FLOAT")
FUNCTION = "filter_image_by_score" FUNCTION = "filter_image"
CATEGORY = "LexTools/ImageProcessing/Scores" CATEGORY = "LexTools/ImageProcessing/Filtering"
def filter_image_by_score(self, score, threshold, image):
# If score > threshold, return the image, otherwise return None
if score < threshold:
pass
else:
return (image,)
def filter_image(self, image, score, threshold, show_on_node):
try:
if float(score) >= float(threshold):
score_text = f"Score {score:.2f} >= Threshold {threshold:.2f}\nImage Passed"
output_ui = {"text": [score_text]} if show_on_node else {}
return {"result": (image, float(score)), "ui": output_ui}
else:
score_text = f"Score {score:.2f} < Threshold {threshold:.2f}\nImage Filtered"
output_ui = {"text": [score_text]} if show_on_node else {}
return {"result": (torch.zeros_like(image), float(score)), "ui": output_ui}
except Exception as e:
print(f"Error filtering image: {str(e)}")
return {"result": (image, 0.0), "ui": {"text": [str(e)]} if show_on_node else {}}
class ImageQualityScoreNode: class ImageQualityScoreNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"aesthetic_score": ("INT", {"default": None}), "aesthetic_score": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0}),
"ai_score_artificial": ("FLOAT", {"default": None}), "image_score_good": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0}),
"ai_score_human": ("FLOAT", {"default": None}), "image_score_bad": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0}),
"show_on_node": ("INT", {"default": 0}), "ai_score_artificial": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
}, "ai_score_human": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0}),
"optional": { "weight_good_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"image_score_good": ("FLOAT", {"default": 0}), "weight_aesthetic_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"image_score_bad": ("FLOAT", {"default": 0}), "weight_bad_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"weight_good_score": ("FLOAT", {"default": 1}), "weight_AIDetection": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"weight_aesthetic_score": ("FLOAT", {"default": 1.0}), "weight_HumanDetection": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
"weight_bad_score": ("FLOAT", {"default": 1.0}), "MultiplyScoreBy": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0}),
"weight_AIDetection": ("FLOAT", {"default": 1.0}), "show_on_node": ("BOOLEAN", {"default": False}),
"weight_HumanDetection": ("FLOAT", {"default": 1.0}),
"MultiplyScoreBy": ("FLOAT", {"default": 100000}),
}, },
} }
OUTPUT_NODE = True
RETURN_TYPES = ("FLOAT",) RETURN_TYPES = ("FLOAT",)
FUNCTION = "calculate_score" FUNCTION = "calculate_score"
CATEGORY = "LexTools/ImageProcessing/Scores" CATEGORY = "LexTools/ImageProcessing/Scoring"
def calculate_score(self, image_score_good, image_score_bad, aesthetic_score, ai_score_artificial, ai_score_human, weight_good_score, weight_aesthetic_score, weight_bad_score, weight_AIDetection, MultiplyScoreBy, show_on_node, weight_HumanDetection): def calculate_score(self, aesthetic_score, image_score_good, image_score_bad, ai_score_artificial, ai_score_human,
# Define the weights and maximum possible values weight_good_score, weight_aesthetic_score, weight_bad_score, weight_AIDetection, weight_HumanDetection,
maxA, maxB, maxC = 3, 3, 1000 MultiplyScoreBy, show_on_node):
# Compute the exponential effect of the AI score try:
ai_score_artificial_exp = 10 ** ai_score_artificial # Calculate weighted scores
# Compute the final score according to the provided formula weighted_aesthetic = float(aesthetic_score) * weight_aesthetic_score
final_score = ((((((image_score_good + maxA) / (2 * maxA) * weight_good_score) + (aesthetic_score / maxC) * weight_bad_score) / (weight_good_score + weight_bad_score)) - weighted_good = float(image_score_good) * weight_good_score
weight_aesthetic_score * ((image_score_bad + maxB) / (2 * maxB))) * ((weight_HumanDetection * (ai_score_human))-(weight_AIDetection * (ai_score_artificial_exp)))) * MultiplyScoreBy weighted_bad = float(image_score_bad) * weight_bad_score
weighted_ai = float(ai_score_artificial) * weight_AIDetection
weighted_human = float(ai_score_human) * weight_HumanDetection
# Prepare the output UI # Calculate total score
return (final_score, {"ui": {"STRING": [final_score]}}) total_score = (weighted_aesthetic + weighted_good - weighted_bad + weighted_human - weighted_ai) * MultiplyScoreBy
# Format score for display
score_text = f"Score: {total_score:.2f}\n"
score_text += f"Aesthetic (w:{weight_aesthetic_score:.1f}): {aesthetic_score:.2f}\n"
score_text += f"Good (w:{weight_good_score:.1f}): {image_score_good:.2f}\n"
score_text += f"Bad (w:{weight_bad_score:.1f}): {image_score_bad:.2f}\n"
score_text += f"AI (w:{weight_AIDetection:.1f}): {ai_score_artificial:.2f}\n"
score_text += f"Human (w:{weight_HumanDetection:.1f}): {ai_score_human:.2f}\n"
score_text += f"Multiplier: {MultiplyScoreBy:.1f}"
output_ui = {"text": [score_text]} if show_on_node else {}
return {"result": (float(total_score),), "ui": output_ui}
except Exception as e:
print(f"Error calculating score: {str(e)}")
return {"result": (0.0,), "ui": {"text": [str(e)]} if show_on_node else {}}
# #
@@ -468,155 +331,59 @@ class MLP(pl.LightningModule):
self.ycol = ycol self.ycol = ycol
self.layers = nn.Sequential( self.layers = nn.Sequential(
nn.Linear(self.input_size, 1024), nn.Linear(self.input_size, 1024),
# nn.ReLU(), #nn.ReLU(),
nn.Dropout(0.2), nn.Dropout(0.2),
nn.Linear(1024, 128), nn.Linear(1024, 128),
# nn.ReLU(), #nn.ReLU(),
nn.Dropout(0.2), nn.Dropout(0.2),
nn.Linear(128, 64), nn.Linear(128, 64),
# nn.ReLU(), #nn.ReLU(),
nn.Dropout(0.1), nn.Dropout(0.1),
nn.Linear(64, 16), nn.Linear(64, 16),
# nn.ReLU(), #nn.ReLU(),
nn.Linear(16, 1) nn.Linear(16, 1)
) )
def forward(self, x): def forward(self, x):
return self.layers(x) return self.layers(x)
def training_step(self, batch, batch_idx): def training_step(self, batch, batch_idx):
x = batch[self.xcol] x = batch[self.xcol]
y = batch[self.ycol].reshape(-1, 1) y = batch[self.ycol].reshape(-1, 1)
x_hat = self.layers(x) x_hat = self.layers(x)
loss = F.mse_loss(x_hat, y) loss = F.mse_loss(x_hat, y)
return loss return loss
def validation_step(self, batch, batch_idx): def validation_step(self, batch, batch_idx):
x = batch[self.xcol] x = batch[self.xcol]
y = batch[self.ycol].reshape(-1, 1) y = batch[self.ycol].reshape(-1, 1)
x_hat = fastapiself.layers(x) x_hat =fastapiself.layers(x)
loss = fastapi.mse_loss(x_hat, y) loss = fastapi.mse_loss(x_hat, y)
return loss return loss
def configure_optimizers(self): def configure_optimizers(self):
optimizer = torch.optim.Adam(self.parameters(), lr=1e-3) optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)
return optimizer return optimizer
def normalized(a, axis=-1, order=2):
def normalized(a, axis=-1, order=2): import numpy as np # pylint: disable=import-outside-toplevel
import numpy as np # pylint: disable=import-outside-toplevel l2 = np.atleast_1d(np.linalg.norm(a, order, axis))
l2 = np.atleast_1d(np.linalg.norm(a, order, axis)) l2[l2 == 0] = 1
l2[l2 == 0] = 1 return a / np.expand_dims(l2, axis)
return a / np.expand_dims(l2, axis)
class AesteticModel: class AesteticModel:
def __init__(self): def __init__(self):
pass pass
@classmethod
@classmethod def INPUT_TYPES(s):
def INPUT_TYPES(s): return { "required": {"model_name": (folder_paths.get_filename_list("aesthetic"), )}}
return {"required": {"model_name": (folder_paths.get_filename_list("aesthetic"), )}} RETURN_TYPES = ("AESTHETIC_MODEL",)
RETURN_TYPES = ("AESTHETIC_MODEL",) FUNCTION = "load_model"
FUNCTION = "load_model" CATEGORY = "LexTools/ImageProcessing/aestheticscore"
CATEGORY = "LexTools/ImageProcessing/aestheticscore" def load_model(self, model_name):
#load model
def load_model(self, model_name): m_path = folder_paths.folder_names_and_paths["aesthetic"][0]
# load model m_path2 = os.path.join(m_path[0],model_name)
m_path = folder_paths.folder_names_and_paths["aesthetic"][0] return (m_path2,)
m_path2 = os.path.join(m_path[0], model_name)
return (m_path2,)
class SaturationMatchingNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"working_image": ("IMAGE",),
"master_image": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "run"
CATEGORY = "LexTools/ImageProcessing/SaturationMatching"
def __init__(self):
self.working_image = None
self.master_image = None
def calculate_saturation(self, tensor_image):
if len(tensor_image.shape) == 4:
# Take the first image from the batch
tensor_image = tensor_image[0]
# Convert tensor to numpy array and then to PIL Image
img = (tensor_image * 255).to(torch.uint8).numpy()
pil_image = Image.fromarray(img, mode='RGB')
# Convert to HSV
hsv_image = pil_image.convert('HSV')
s_channel = np.array(hsv_image)[:, :, 1]
# Calculate average saturation
avg_saturation = np.mean(s_channel)
return avg_saturation
def adjust_saturation(self, tensor_image, target_saturation):
original_shape = tensor_image.shape
if len(tensor_image.shape) == 4:
# Take the first image from the batch
tensor_image = tensor_image[0]
# Convert tensor to numpy array and then to PIL Image
img = (tensor_image * 255).to(torch.uint8).cpu().numpy()
img = np.transpose(img, (1, 2, 0))
pil_image = Image.fromarray(img, mode='RGB')
# Convert to HSV
hsv_image = pil_image.convert('HSV')
hsv_array = np.array(hsv_image)
# Calculate current average saturation
current_saturation = np.mean(hsv_array[:, :, 1])
# Calculate adjustment factor
factor = target_saturation / current_saturation
# Adjust saturation
hsv_array[:, :, 1] = np.clip(
hsv_array[:, :, 1] * factor, 0, 255).astype(np.uint8)
# Convert back to PIL Image and then to tensor
adjusted_hsv_image = Image.fromarray(hsv_array, 'HSV')
adjusted_rgb_image = adjusted_hsv_image.convert('RGB')
tensor_transform = transforms.ToTensor()
adjusted_tensor = tensor_transform(adjusted_rgb_image)
# Reshape to match the original tensor shape
if len(original_shape) == 4:
adjusted_tensor = adjusted_tensor.unsqueeze(0)
return adjusted_tensor
def run(self, working_image, master_image):
self.working_image = working_image
self.master_image = master_image
# Calculate target saturation from master image
target_saturation = self.calculate_saturation(self.master_image)
# Adjust the saturation of the working image
adjusted_image = self.adjust_saturation(
self.working_image, target_saturation)
return adjusted_image
class CalculateAestheticScore: class CalculateAestheticScore:
device = "cuda" device = "cuda"
model2 = None model2 = None
preprocess = None preprocess = None
model = None model = None
@@ -632,7 +399,7 @@ class CalculateAestheticScore:
"aesthetic_model": ("AESTHETIC_MODEL",), "aesthetic_model": ("AESTHETIC_MODEL",),
}, },
"optional": { "optional": {
"keep_in_memory": ("BOOL", {"default": True}), "keep_in_memory": ("BOOLEAN", {"default": True}),
} }
} }
@@ -642,18 +409,16 @@ class CalculateAestheticScore:
def execute(self, image, aesthetic_model, keep_in_memory): def execute(self, image, aesthetic_model, keep_in_memory):
if not self.model2 or not self.preprocess: if not self.model2 or not self.preprocess:
self.model2, self.preprocess = clip.load( self.model2, self.preprocess = clip.load("ViT-L/14", device=self.device) #RN50x64
"ViT-L/14", device=self.device) # RN50x64
m_path2 = aesthetic_model m_path2 = aesthetic_model
if not self.model: if not self.model:
# CLIP embedding dim is 768 for CLIP ViT L 14 self.model = MLP(768) # CLIP embedding dim is 768 for CLIP ViT L 14
self.model = MLP(768)
s = torch.load(m_path2) s = torch.load(m_path2)
self.model.load_state_dict(s) self.model.load_state_dict(s)
self.model.to(self.device) self.model.to(self.device)
self.model.eval() self.model.eval()
tensor_image = image[0] tensor_image = image[0]
@@ -668,8 +433,7 @@ class CalculateAestheticScore:
image_features = self.model2.encode_image(image2) image_features = self.model2.encode_image(image2)
im_emb_arr = normalized(image_features.cpu().detach().numpy()) im_emb_arr = normalized(image_features.cpu().detach().numpy())
prediction = self.model(torch.from_numpy(im_emb_arr).to( prediction = self.model(torch.from_numpy(im_emb_arr).to(self.device).type(torch.cuda.FloatTensor))
self.device).type(torch.cuda.FloatTensor))
final_prediction = int(float(prediction[0])*100) final_prediction = int(float(prediction[0])*100)
if not keep_in_memory: if not keep_in_memory:
@@ -678,11 +442,10 @@ class CalculateAestheticScore:
self.preprocess = None self.preprocess = None
return (final_prediction,) return (final_prediction,)
class MD5ImageHashNode: class MD5ImageHashNode:
device = "cuda" device = "cuda"
def __init__(self): def __init__(self):
pass pass
@@ -700,7 +463,7 @@ class MD5ImageHashNode:
def execute(self, image): def execute(self, image):
tensor_image = image[0] tensor_image = image[0]
# Convert the tensor to a PIL image # Convert the tensor to a PIL image
img = (tensor_image * 255).to(torch.uint8).cpu().numpy() img = (tensor_image * 255).to(torch.uint8).cpu().numpy()
pil_image = Image.fromarray(img, mode='RGB') pil_image = Image.fromarray(img, mode='RGB')
@@ -717,33 +480,29 @@ class MD5ImageHashNode:
return (md5_hash,) return (md5_hash,)
class AesthetlcScoreSorter: class AesthetlcScoreSorter:
def __init__(self): def __init__(self):
pass
pass pass
pass
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required":{
"image": ("IMAGE",), "image": ("IMAGE",),
"score": ("SCORE",), "score": ("SCORE",),
"image2": ("IMAGE",), "image2": ("IMAGE",),
"score2": ("SCORE",), "score2": ("SCORE",),
} }
} }
RETURN_TYPES = ("IMAGE", "SCORE", "IMAGE", "SCORE",) RETURN_TYPES = ("IMAGE", "SCORE", "IMAGE", "SCORE",)
FUNCTION = "execute" FUNCTION = "execute"
CATEGORY = "LexTools/ImageProcessing/aestheticscore" CATEGORY = "LexTools/ImageProcessing/aestheticscore"
def execute(self,image,score,image2,score2):
def execute(self, image, score, image2, score2): if score >= score2:
if score >= score2: return (image, score, image2, score2,)
return (image, score, image2, score2,) else:
else: return (image2, score2, image, score,)
return (image2, score2, image, score,)
class ScoreConverterNode: class ScoreConverterNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
@@ -769,36 +528,34 @@ class ScoreConverterNode:
# Prepare the output UI # Prepare the output UI
output_ui = {} output_ui = {}
if show_on_node == 1: if show_on_node ==1:
output_ui = {"ui": {"STRING": [score_str]}} output_ui = {"ui": {"STRING": [score_str]}}
return (score_int, score_float, score_str, output_ui)
return (score_int, score_float, score_str, output_ui)
class SamplerPropertiesNode: class SamplerPropertiesNode:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {"required":{
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}), "steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.5, "round": 0.01}), "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.5, "round": 0.01}),
"denoise": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 1, "step": 0.1, "round": 0.01}), "denoise": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 1, "step":0.1, "round": 0.01}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, ), "sampler_name": (comfy.samplers.KSampler.SAMPLERS, ),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
}
}
} RETURN_TYPES = ("STRING","INT","FLOAT","FLOAT","STRING","STRING")
}
RETURN_TYPES = ("STRING", "INT", "FLOAT", "FLOAT", "STRING", "STRING")
FUNCTION = "sample" FUNCTION = "sample"
CATEGORY = "sampling" CATEGORY = "sampling"
def sample(self, ckpt_name, steps, cfg, sampler_name, scheduler, denoise): def sample(self, ckpt_name, steps, cfg, sampler_name, scheduler,denoise):
pass pass
return (ckpt_name, steps, cfg, sampler_name, scheduler, denoise) return (ckpt_name, steps, cfg, sampler_name, scheduler,denoise)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"ImageFilterByIntScoreNode": ImageFilterByIntScoreNode, "ImageFilterByIntScoreNode": ImageFilterByIntScoreNode,
@@ -807,11 +564,12 @@ NODE_CLASS_MAPPINGS = {
"ImageAspectPadNode": ImageAspectPadNode, "ImageAspectPadNode": ImageAspectPadNode,
"ImageRankingNode": ImageRankingNode, "ImageRankingNode": ImageRankingNode,
"ImageQualityScoreNode": ImageQualityScoreNode, "ImageQualityScoreNode": ImageQualityScoreNode,
"ScoreConverterNode": ScoreConverterNode, "ScoreConverterNode":ScoreConverterNode,
"MD5ImageHashNode": MD5ImageHashNode, "MD5ImageHashNode": MD5ImageHashNode,
"SamplerPropertiesNode": SamplerPropertiesNode, "SamplerPropertiesNode": SamplerPropertiesNode,
"CalculateAestheticScore": CalculateAestheticScore,
"SaturationMatchingNode": SaturationMatchingNode "LoadAesteticModel":AesteticModel,
"AesthetlcScoreSorter": AesthetlcScoreSorter,
} }
@@ -820,9 +578,10 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ImageFilterByFloatScoreNode": "Image Filter (Float Score)", "ImageFilterByFloatScoreNode": "Image Filter (Float Score)",
"ImageScaleToMin": "Image Scale To Min", "ImageScaleToMin": "Image Scale To Min",
"ImageRankingNode": "Image Ranking For Image Reward", "ImageRankingNode": "Image Ranking For Image Reward",
"ScoreConverterNode": "Score Converter (Aesthetic Score)", "ScoreConverterNode":"Score Converter (Aesthetic Score)",
"MD5ImageHashNode": "MD5 Image Hash", "MD5ImageHashNode":"MD5 Image Hash",
"SamplerPropertiesNode": "Property Output Node.", "SamplerPropertiesNode":"Sampler input node",
"SaturationMatchingNode": "SaturationMatchingNode", "LoadAesteticModel": "LoadAesteticModel",
"CalculateAestheticScore": "CalculateAestheticScore",
} "AesthetlcScoreSorter": "AesthetlcScoreSorter",
}
+435 -135
View File
@@ -6,35 +6,99 @@ import matplotlib.pyplot as plt
import numpy as np import numpy as np
import io import io
from scipy.ndimage import binary_dilation from scipy.ndimage import binary_dilation
import os
from pathlib import Path
import json
model_names = [ model_names = [
"enes361/segformer_b2_clothes", "sayeed99/segformer_b3_clothes",
"mattmdjaga/segformer_b0_clothes", "mattmdjaga/segformer_b0_clothes",
"mattmdjaga/segformer_b2_clothes", "mattmdjaga/segformer_b2_clothes",
"DiTo97/binarization-segformer-b3", "DiTo97/binarization-segformer-b3",
"s3nh/SegFormer-b0-person-segmentation", "s3nh/SegFormer-b0-person-segmentation",
"venture361/clothes_segmentation", "venture361/clothes_segmentation",
"itsitgroup/human-body-segmentation",
"matei-dorian/segformer-b5-finetuned-human-parsing", "matei-dorian/segformer-b5-finetuned-human-parsing",
"Lexic0n/segformer-b0-finetuned-human-parsing", "Lexic0n/segformer-b0-finetuned-human-parsing",
"sam1120/segformer-b0-finetuned-neurosymbolic-contingency-bag1-v0.1-v0", "sam1120/segformer-b0-finetuned-neurosymbolic-contingency-bag1-v0.1-v0",
"ehsanhallo/segformer-b0-scene-parse-150" "ehsanhallo/segformer-b0-scene-parse-150"
] ]
class SegformerModelLoader:
_models = {} # Cache for loaded models
_processors = {} # Cache for loaded processors
@classmethod
def get_local_checkpoints(cls):
"""Get list of local checkpoint directories"""
checkpoints_dir = Path("models/segformer")
if not checkpoints_dir.exists():
checkpoints_dir.mkdir(parents=True, exist_ok=True)
return []
# Look for config.json files in subdirectories
checkpoints = []
for path in checkpoints_dir.glob("*/config.json"):
checkpoints.append(path.parent.name)
return checkpoints
@classmethod
def load_model(cls, model_name, local_dir=None):
"""Load model and processor with caching"""
# Check cache first
cache_key = model_name if not local_dir else str(local_dir)
if cache_key in cls._models:
return cls._models[cache_key], cls._processors[cache_key]
try:
if local_dir:
processor = SegformerImageProcessor.from_pretrained(local_dir)
model = AutoModelForSemanticSegmentation.from_pretrained(local_dir)
else:
processor = SegformerImageProcessor.from_pretrained(model_name)
model = AutoModelForSemanticSegmentation.from_pretrained(model_name)
# Cache the loaded model and processor
cls._models[cache_key] = model
cls._processors[cache_key] = processor
return model, processor
except Exception as e:
print(f"Error loading model {model_name}: {str(e)}")
# Fallback to a reliable model
return cls.load_model("matei-dorian/segformer-b5-finetuned-human-parsing")
@classmethod
def clear_cache(cls):
"""Clear the model cache"""
cls._models.clear()
cls._processors.clear()
# Update the model_names list to include local checkpoints
def get_available_models():
local_checkpoints = SegformerModelLoader.get_local_checkpoints()
return model_names + [f"local:{cp}" for cp in local_checkpoints]
class SegformerNode: class SegformerNode:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
global model_names # Assuming model_names is a list of model names
return { return {
"required": { "required": {
"image": ("IMAGE", {"default": None}), "image": ("IMAGE", {"default": None}),
"model_name": (model_names, {"default": model_names[0]}), "model_name": (get_available_models(), {"default": model_names[0]}),
"normalize_mask": ("BOOLEAN", {"default": True}),
"binary_mask": ("BOOLEAN", {"default": False}),
"resize_mode": (["nearest", "bilinear", "bicubic"], {"default": "bilinear"}),
"invert_mask": ("BOOLEAN", {"default": False}),
"show_preview": ("BOOLEAN", {"default": True}),
"return_individual_masks": ("BOOLEAN", {"default": False}),
"post_process": (["none", "erode", "dilate", "smooth"], {"default": "none"}),
"post_process_radius": ("INT", {"default": 3, "min": 1, "max": 10}),
"segment_groups": ("STRING", {"default": "", "multiline": True}),
}, },
} }
RETURN_TYPES = ("IMAGE","MASK", "STRING") RETURN_TYPES = ("IMAGE", "MASK", "STRING", "IMAGE") # Added IMAGE for preview
FUNCTION = "segment_image" FUNCTION = "segment_image"
CATEGORY = "LexTools/ImageProcessing/Segmentation" CATEGORY = "LexTools/ImageProcessing/Segmentation"
@@ -43,79 +107,199 @@ class SegformerNode:
# self.processor = SegformerImageProcessor.from_pretrained("mattmdjaga/segformer_b2_clothes") # self.processor = SegformerImageProcessor.from_pretrained("mattmdjaga/segformer_b2_clothes")
# self.model = AutoModelForSemanticSegmentation.from_pretrained("mattmdjaga/segformer_b2_clothes") # self.model = AutoModelForSemanticSegmentation.from_pretrained("mattmdjaga/segformer_b2_clothes")
def segment_image(self, image,model_name,): def process_mask(self, mask, normalize=True, binary=False, invert=False, post_process="none", radius=3):
show_on_node = False # Convert to float32 if not already
self.processor = SegformerImageProcessor.from_pretrained(model_name) mask = mask.float()
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name)
i = 255. * image[0].cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
inputs = self.processor(images=img, return_tensors="pt")
outputs = self.model(**inputs)
logits = outputs.logits.cpu()
upsampled_logits = nn.functional.interpolate(
logits,
size=img.size[::-1],
mode="bilinear",
align_corners=False,
)
pred_seg = upsampled_logits.argmax(dim=1)[0]
# Convert the matplotlib figure to a PIL Image and return it
fig = plt.figure()
plt.imshow(pred_seg)
buf = io.BytesIO()
plt.savefig(buf, format='png')
buf.seek(0)
img2 = Image.open(buf)
i = ImageOps.exif_transpose(img2) # Normalize to 0-1 range if requested
if i.getbands() != ("R", "G", "B", "A"): if normalize:
i = i.convert("RGBA") mask = (mask - mask.min()) / (mask.max() - mask.min() + 1e-8)
# Convert to binary if requested
if binary:
mask = (mask > 0.5).float()
# Apply post-processing
if post_process != "none":
kernel = torch.ones(2 * radius + 1, 2 * radius + 1)
if post_process == "erode":
mask = torch.nn.functional.conv2d(
mask.unsqueeze(0).unsqueeze(0),
kernel.unsqueeze(0).unsqueeze(0),
padding=radius
).squeeze() < kernel.sum()
elif post_process == "dilate":
mask = torch.nn.functional.conv2d(
mask.unsqueeze(0).unsqueeze(0),
kernel.unsqueeze(0).unsqueeze(0),
padding=radius
).squeeze() > 0
elif post_process == "smooth":
mask = torch.nn.functional.conv2d(
mask.unsqueeze(0).unsqueeze(0),
kernel.unsqueeze(0).unsqueeze(0),
padding=radius
).squeeze() / kernel.sum()
mask = mask.float()
# Invert if requested
if invert:
mask = 1 - mask
img2 = np.array(img2).astype(np.float32) / 255.0 return mask
img2 = torch.from_numpy(img2)[None,]
if 'A' in i.getbands(): def create_preview(self, image, mask):
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 # Ensure image is in CHW format
mask = 1. - torch.from_numpy(mask) if len(image.shape) == 2:
else: image = image.unsqueeze(0).repeat(3, 1, 1)
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") elif len(image.shape) == 3:
# Get the unique segments in the image if image.shape[0] != 3: # If channels are not in first dimension
unique_segments = np.unique(pred_seg) image = image.permute(2, 0, 1) # Move channels to first dimension
# Ensure mask has correct dimensions
if len(mask.shape) == 3:
mask = mask.squeeze(0)
if len(mask.shape) > 2:
mask = mask.squeeze()
# Resize mask to match image dimensions if needed
if mask.shape != image.shape[1:]:
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0),
size=image.shape[1:],
mode='bilinear',
align_corners=False
).squeeze()
# Create preview by concatenating image and mask
preview = image.clone()
preview = torch.cat([preview, mask.unsqueeze(0)], dim=0)
return preview
# Create a string with the information for each segment def parse_segment_groups(self, groups_str):
segment_info = [] if not groups_str.strip():
for segment in unique_segments: return {}
# Get the name of the segment from the model's configuration
segment_name = self.model.config.id2label[segment] groups = {}
for line in groups_str.split('\n'):
if ':' in line:
name, indices = line.split(':')
indices = [int(i.strip()) for i in indices.split(',') if i.strip()]
groups[name.strip()] = indices
return groups
# Here, you would replace these values with the actual accuracy and IoU for the segment def segment_image(self, image, model_name, normalize_mask=True, binary_mask=False,
resize_mode="bilinear", invert_mask=False, show_preview=True,
return_individual_masks=False, post_process="none",
post_process_radius=3, segment_groups=""):
try:
# Handle local checkpoint loading
if model_name.startswith("local:"):
local_dir = Path("models/segformer") / model_name[6:]
self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir)
else:
self.model, self.processor = SegformerModelLoader.load_model(model_name)
show_on_node = False
# Process input image
i = 255. * image[0].cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
inputs = self.processor(images=img, return_tensors="pt")
segment_info.append(f"Segment {segment}: {segment_name}") # Get model outputs
outputs = self.model(**inputs)
logits = outputs.logits.cpu()
# Join the segment info strings into a single string # Upsample logits with specified resize mode
segment_info_str = "\n".join(segment_info) upsampled_logits = nn.functional.interpolate(
logits,
size=img.size[::-1],
mode=resize_mode,
align_corners=False if resize_mode != "nearest" else None,
)
pred_seg = upsampled_logits.argmax(dim=1)[0]
# Parse segment groups if provided
segment_groups_dict = self.parse_segment_groups(segment_groups)
# Create individual masks if requested
individual_masks = {}
segment_info = []
# Get unique segments and process each
unique_segments = np.unique(pred_seg.numpy())
for segment in unique_segments:
segment_name = self.model.config.id2label[segment]
segment_info.append(f"Segment {segment}: {segment_name}")
if return_individual_masks:
mask = (pred_seg == segment).float()
mask = self.process_mask(mask, normalize_mask, binary_mask,
invert_mask, post_process, post_process_radius)
individual_masks[segment_name] = mask
output_ui = {"images": [img2]} if show_on_node else {} # Create merged mask based on segment groups
if segment_groups_dict:
merged_mask = torch.zeros_like(pred_seg, dtype=torch.float32)
for group_name, indices in segment_groups_dict.items():
group_mask = torch.zeros_like(pred_seg, dtype=torch.float32)
for idx in indices:
group_mask = torch.maximum(group_mask, (pred_seg == idx).float())
merged_mask = torch.maximum(merged_mask, group_mask)
segment_info.append(f"Group {group_name}: {indices}")
else:
merged_mask = torch.ones_like(pred_seg, dtype=torch.float32)
return {"result": (img2,mask, segment_info_str), "ui": output_ui} # Process the final mask
merged_mask = self.process_mask(merged_mask, normalize_mask, binary_mask,
invert_mask, post_process, post_process_radius)
# Create visualization
fig = plt.figure()
plt.imshow(pred_seg)
buf = io.BytesIO()
plt.savefig(buf, format='png')
buf.seek(0)
img2 = Image.open(buf)
# Convert visualization to tensor
i = ImageOps.exif_transpose(img2)
if i.getbands() != ("R", "G", "B", "A"):
i = i.convert("RGBA")
img2 = np.array(img2).astype(np.float32) / 255.0
img2 = torch.from_numpy(img2)[None,]
# Create preview if requested
preview = self.create_preview(image[0], merged_mask) if show_preview else None
# Join segment info
segment_info_str = "\n".join(segment_info)
if return_individual_masks:
segment_info_str += "\n\nIndividual masks available for: " + ", ".join(individual_masks.keys())
output_ui = {"images": [img2]} if show_on_node else {}
# Return results
return {"result": (img2, merged_mask, segment_info_str, preview if preview is not None else img2),
"ui": output_ui}
except Exception as e:
import traceback
print(f"Error in segmentation: {str(e)}")
print(f"Traceback: {traceback.format_exc()}")
return {"result": (image, torch.zeros_like(image[0, :, :]), f"Error: {str(e)}", image),
"ui": {"images": [image]} if show_on_node else {}}
class SegformerNodeMasks: class SegformerNodeMasks:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
global model_names # Assuming model_names is a list of model names
return { return {
"required": { "required": {
"image": ("IMAGE", {"default": None}), "image": ("IMAGE", {"default": None}),
"segments_to_merge": ("STRING", {"default": "0"}), "segments_to_merge": ("STRING", {"default": "0"}),
"model_name": (model_names, {"default": model_names[0]}) "model_name": (get_available_models(), {"default": model_names[0]}),
}, },
} }
@@ -128,14 +312,17 @@ class SegformerNodeMasks:
# Function to segment the image and return the merged segments as per the provided indices # Function to segment the image and return the merged segments as per the provided indices
def segment_image(self, image, segments_to_merge, model_name): def segment_image(self, image, segments_to_merge, model_name):
# Handle local checkpoint loading
if model_name.startswith("local:"):
local_dir = Path("models/segformer") / model_name[6:]
self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir)
else:
self.model, self.processor = SegformerModelLoader.load_model(model_name)
# Convert the segments_to_merge from string to list of integers # Convert the segments_to_merge from string to list of integers
show_on_node=False show_on_node=False
segments_to_merge = list(map(int, segments_to_merge.split(','))) segments_to_merge = list(map(int, segments_to_merge.split(',')))
# Load the pretrained models and processors
self.processor = SegformerImageProcessor.from_pretrained(model_name)
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name)
# Preprocess the image # Preprocess the image
i = 255. * image[0].cpu().numpy() i = 255. * image[0].cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
@@ -195,106 +382,219 @@ class SegformerNodeMasks:
merged_image_pil = Image.fromarray(merged_image) merged_image_pil = Image.fromarray(merged_image)
img2 = torch.from_numpy(np.array(merged_image_pil).astype(np.float32) / 255.0)[None,] img2 = torch.from_numpy(np.array(merged_image_pil).astype(np.float32) / 255.0)[None,]
# Convert the merged mask to byte format (0-255) and ensure correct dimensionality
merged_mask = (merged_mask > 0).float() # Convert to binary mask first
merged_mask = torch.clamp(merged_mask, 0, 1)
output_ui = {"images": [img2]} if show_on_node else {} output_ui = {"images": [img2]} if show_on_node else {}
return {"result": (img2, merged_mask, 'Merged Segments'), "ui": output_ui} return {"result": (img2, merged_mask, 'Merged Segments'), "ui": output_ui}
class SegformerNodeMergeSegments: class SegformerNodeMergeSegments:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
global model_names
return { return {
"required": { "required": {
"image": ("IMAGE", {"default": None}), "image": ("IMAGE", {"default": None}),
"segments_to_merge_str": ("STRING", {"default": ""}), "segments_to_merge_str": ("STRING", {"default": ""}),
"model_name": (model_names, {"default": model_names[0]}), "model_name": (get_available_models(), {"default": model_names[0]}),
"blur_radius": ("INT", {"default": 0}), "normalize_mask": ("BOOLEAN", {"default": True}),
"dilation_radius": ("INT", {"default": 0}), # Added dilation_radius "binary_mask": ("BOOLEAN", {"default": False}),
"intensity": ("FLOAT", {"default": 1.0}), # Added intensity "resize_mode": (["nearest", "bilinear", "bicubic"], {"default": "bilinear"}),
"ceiling": ("FLOAT", {"default": 1.0}), # Added ceiling "invert_mask": ("BOOLEAN", {"default": False}),
"show_preview": ("BOOLEAN", {"default": True}),
"blur_radius": ("INT", {"default": 5, "min": 0, "max": 100}),
"dilation_radius": ("INT", {"default": 5, "min": 0, "max": 100}),
"intensity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
"ceiling": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}),
}, },
} }
OUTPUT_NODE = True OUTPUT_NODE = True
RETURN_TYPES = ("IMAGE", "MASK", "STRING", "IMAGE") # Added IMAGE for preview
RETURN_TYPES = ("IMAGE", "MASK", "STRING")
FUNCTION = "merge_segments" FUNCTION = "merge_segments"
CATEGORY = "LexTools/ImageProcessing/Segmentation" CATEGORY = "LexTools/ImageProcessing/Segmentation"
def __init__(self): def __init__(self):
pass pass
def merge_segments(self, image, segments_to_merge_str, model_name, blur_radius, dilation_radius, intensity, ceiling): # Added dilation_radius in the arguments def process_mask(self, mask, normalize=True, binary=False, invert=False, blur_radius=0, dilation_radius=0, intensity=1.0, ceiling=1.0):
# Convert to float32 if not already
show_on_node=False if isinstance(mask, np.ndarray):
mask = torch.from_numpy(mask)
mask = mask.float()
# Ensure mask is 2D
if len(mask.shape) > 2:
mask = mask.squeeze()
# Normalize to 0-1 range if requested
if normalize:
min_val = mask.min()
max_val = mask.max()
if max_val > min_val:
mask = (mask - min_val) / (max_val - min_val)
# Convert to binary if requested
if binary:
mask = (mask > 0.5).float()
# Apply dilation if specified
if dilation_radius > 0:
kernel = torch.ones(2 * dilation_radius + 1, 2 * dilation_radius + 1)
mask = torch.nn.functional.conv2d(
mask.unsqueeze(0).unsqueeze(0),
kernel.unsqueeze(0).unsqueeze(0),
padding=dilation_radius
).squeeze() > 0
mask = mask.float()
# Apply Gaussian blur for feathering
if blur_radius > 0:
# Ensure mask is 2D and in correct range for PIL
mask_np = (mask.squeeze().numpy() * 255).astype(np.uint8)
mask_pil = Image.fromarray(mask_np, mode='L') # Use 'L' mode for grayscale
mask_pil = mask_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius))
mask = torch.from_numpy(np.array(mask_pil).astype(np.float32) / 255.0)
# Apply intensity and ceiling
mask = torch.clamp(mask * intensity, 0, ceiling)
# Invert if requested
if invert:
mask = 1 - mask
return mask
def create_preview(self, image, mask):
# Create an RGBA preview with the mask as alpha channel
if len(image.shape) == 2:
image = image.unsqueeze(0).repeat(3, 1, 1)
elif len(image.shape) == 3:
if image.shape[0] != 3: # If channels are not in first dimension
image = image.permute(2, 0, 1) # Move channels to first dimension
# Ensure mask has correct dimensions
if len(mask.shape) == 3:
mask = mask.squeeze(0)
if len(mask.shape) > 2:
mask = mask.squeeze()
if mask.shape != image.shape[1:]:
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0),
size=image.shape[1:],
mode='bilinear',
align_corners=False
).squeeze()
preview = image.clone()
preview = torch.cat([preview, mask.unsqueeze(0)], dim=0)
return preview
def merge_segments(self, image, segments_to_merge_str, model_name, normalize_mask=True,
binary_mask=False, resize_mode="bilinear", invert_mask=False,
show_preview=True, blur_radius=5, dilation_radius=5,
intensity=1.0, ceiling=1.0):
try: try:
self.processor = SegformerImageProcessor.from_pretrained(model_name) # Handle local checkpoint loading
except Exception: if model_name.startswith("local:"):
print(f"Failed to load preprocessor for model {model_name}. Using preprocessor from mattmdjaga/segformer_b2_clothes instead.") local_dir = Path("models/segformer") / model_name[6:]
self.processor = SegformerImageProcessor.from_pretrained("matei-dorian/segformer-b5-finetuned-human-parsing") self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir)
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name) else:
self.model, self.processor = SegformerModelLoader.load_model(model_name)
show_on_node = False
i = 255. * image[0].cpu().numpy() # Get input image dimensions and ensure proper shape
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) input_image = image[0].cpu()
inputs = self.processor(images=img, return_tensors="pt") if len(input_image.shape) != 3:
raise ValueError(f"Expected input image with shape (H,W,C) or (C,H,W), got {input_image.shape}")
outputs = self.model(**inputs)
logits = outputs.logits.cpu()
upsampled_logits = nn.functional.interpolate(
logits,
size=img.size[::-1],
mode="bilinear",
align_corners=False,
)
pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
unique_segments = np.unique(pred_seg)
segments_to_merge = list(map(int, segments_to_merge_str.split(',')))
merged_mask = np.zeros_like(pred_seg)
merged_segments = []
for segment in unique_segments:
if segment in segments_to_merge:
mask = np.where(pred_seg == segment, 1, 0)
mask = nn.functional.interpolate(torch.from_numpy(mask.astype(np.float32))[None, None,], size=(img.height, img.width), mode="nearest")[0,0].numpy()
merged_mask = np.maximum(merged_mask, mask)
merged_segments.append(segment)
merged_mask = np.clip(merged_mask * intensity, 0, ceiling) # Apply intensity and ceiling to the mask
if dilation_radius > 0: # Dilate the mask if dilation_radius > 0
struct = np.ones((2 * dilation_radius + 1, 2 * dilation_radius + 1))
merged_mask = binary_dilation(merged_mask, structure=struct)
merged_mask_rgb = np.repeat(merged_mask[..., None], 3, axis=2)
if blur_radius > 0: # Blur the mask if radius > 0
merged_mask_rgb = Image.fromarray((merged_mask_rgb * 255).astype('uint8'))
merged_mask_rgb = merged_mask_rgb.filter(ImageFilter.GaussianBlur(radius=blur_radius))
merged_mask_rgb = np.array(merged_mask_rgb) / 255.0
merged_image = np.array(img) * merged_mask_rgb
merged_image_pil = Image.fromarray(merged_image.astype('uint8'))
if blur_radius > 0: # Apply blur if radius > 0
merged_image_pil = merged_image_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius))
img2 = np.array(merged_image_pil).astype(np.float32) / 255.0 # Ensure image is in HWC format
img2 = torch.from_numpy(img2).double()[None,] if input_image.shape[0] == 3: # If in CHW format
input_image = input_image.permute(1, 2, 0)
input_height, input_width = input_image.shape[0:2]
# Process input image
img = Image.fromarray((input_image.numpy() * 255).astype(np.uint8))
inputs = self.processor(images=img, return_tensors="pt")
merged_mask_torch = torch.from_numpy(merged_mask).float()[None,] # change from double to float outputs = self.model(**inputs)
logits = outputs.logits.cpu()
merged_segments_str = ','.join(map(str, merged_segments)) # Upsample logits to match input image size
upsampled_logits = nn.functional.interpolate(
logits,
size=(input_height, input_width),
mode=resize_mode,
align_corners=False if resize_mode != "nearest" else None,
)
output_ui = {"images": [img2]} if show_on_node else {} pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
unique_segments = np.unique(pred_seg)
return {"result": (img2, merged_mask_torch, merged_segments_str), "ui": output_ui} # Handle empty segments string
if not segments_to_merge_str.strip():
segments_to_merge = []
else:
segments_to_merge = [int(s.strip()) for s in segments_to_merge_str.split(',') if s.strip()]
merged_mask = np.zeros((input_height, input_width), dtype=np.float32)
merged_segments = []
for segment in unique_segments:
if segment in segments_to_merge:
mask = np.where(pred_seg == segment, 1, 0)
merged_mask = np.maximum(merged_mask, mask)
merged_segments.append(segment)
# Convert to tensor and process
merged_mask = torch.from_numpy(merged_mask)
merged_mask = self.process_mask(
merged_mask,
normalize=normalize_mask,
binary=binary_mask,
invert=invert_mask,
blur_radius=blur_radius,
dilation_radius=dilation_radius,
intensity=intensity,
ceiling=ceiling
)
# Ensure mask has correct dimensions for broadcasting
merged_mask_3d = merged_mask.unsqueeze(-1) # Add channel dimension for broadcasting
# Apply mask to image
merged_image = input_image.numpy() * merged_mask_3d.numpy()
# Convert back to tensor in CHW format
merged_image = torch.from_numpy(merged_image).permute(2, 0, 1).unsqueeze(0)
merged_segments_str = ','.join(map(str, merged_segments))
if not merged_segments:
merged_segments_str = "No segments selected"
# Create preview
if show_preview:
preview = self.create_preview(input_image.permute(2, 0, 1), merged_mask)
else:
preview = merged_image
output_ui = {"images": [merged_image]} if show_on_node else {}
return {"result": (merged_image, merged_mask, merged_segments_str, preview),
"ui": output_ui}
except Exception as e:
import traceback
print(f"Error merging segments: {str(e)}")
print(f"Traceback: {traceback.format_exc()}")
# Return original image and empty mask on error
empty_mask = torch.zeros((input_height, input_width), dtype=torch.float32)
return {"result": (image, empty_mask, f"Error: {str(e)}", image),
"ui": {"images": [image]} if show_on_node else {}}
+31
View File
@@ -0,0 +1,31 @@
from importlib.util import find_spec
from os.path import exists
from subprocess import run
is_portable = True if exists("python_embeded") else False
def check_and_install_package(package_name: str, install_name: str = None) -> None:
if find_spec(package_name):
return
print(f"/_\ Installing {package_name}")
package = install_name if install_name else package_name
if is_portable:
command = f".\\python_embeded\\python.exe -s -m pip install {package}"
else:
command = f"pip install {package}"
process = run(command, shell=True, check=True, capture_output=True)
print("/_\ Checking packages")
check_and_install_package("transformers")
check_and_install_package("pillow")
check_and_install_package("matplotlib")
check_and_install_package("numpy")
check_and_install_package("scipy")
check_and_install_package("fastapi")
check_and_install_package("pytorch_lightning")
check_and_install_package("clip", "git+https://github.com/openai/CLIP.git")
+39
View File
@@ -0,0 +1,39 @@
[project]
name = "comfyui-lextools"
description = """
A comprehensive toolkit for ComfyUI that provides advanced image processing, analysis, and AI-powered features:
- Semantic segmentation with multiple pre-trained models and mask processing
- Image classification (age, food, documents, NSFW content, AI detection)
- Image captioning using BLIP
- Image quality scoring and filtering
- Workflow utilities for seed management and image aspect ratio handling
"""
version = "1.0.2"
license = "MIT"
dependencies = [
"torch>=2.0.0",
"transformers>=4.30.0",
"Pillow>=9.0.0",
"matplotlib>=3.0.0",
"numpy>=1.20.0",
"scipy>=1.7.0",
"huggingface_hub>=0.19.0"
]
[project.urls]
Repository = "https://github.com/SOELexicon/ComfyUI-LexTools"
Documentation = "https://github.com/SOELexicon/ComfyUI-LexTools/blob/main/README.md"
[tool.comfy]
PublisherId = "lexicon"
DisplayName = "ComfyUI-LexTools"
Description = "Advanced image processing and AI analysis toolkit for ComfyUI"
Icon = "🛠️"
Tags = [
"image processing",
"segmentation",
"classification",
"captioning",
"workflow",
"utilities"
]
+2 -1
View File
@@ -1,4 +1,5 @@
numpy numpy
opencv-python opencv-python
git+https://github.com/facebookresearch/detectron2.git git+https://github.com/facebookresearch/detectron2.git
pyodbc pyodbc
pytorch_lightning