Compare commits
26
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8fb9b27ea1 | ||
|
|
877dd93b0c | ||
|
|
6f3b85c157 | ||
|
|
12cff98987 | ||
|
|
9d64aad7cc | ||
|
|
024fdf44b4 | ||
|
|
78a0fe1ddb | ||
|
|
e1b49cc738 | ||
|
|
d7ff05266f | ||
|
|
7ae11ac705 | ||
|
|
d1da468459 | ||
|
|
f17a163bff | ||
|
|
75bf61118b | ||
|
|
ae3b49a80e | ||
|
|
2f75f817e4 | ||
|
|
4e4a19185b | ||
|
|
9dbd068a71 | ||
|
|
2dfc5cfbe5 | ||
|
|
438df02026 | ||
|
|
285ed93e76 | ||
|
|
3341bcea6e | ||
|
|
a559d3815d | ||
|
|
533910af92 | ||
|
|
55dccd5943 | ||
|
|
3cff522e9a | ||
|
|
a41973d0c3 |
@@ -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 }}
|
||||
@@ -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
|
||||
@@ -25,65 +25,100 @@ ComfyUI-LexTools is a Python-based image processing and analysis toolkit that us
|
||||
- _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:
|
||||
- `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)
|
||||
- `AesthetlcScoreSorter`: Sorts the images by score. (No specific input or output detailed in the provided code)
|
||||
- `AesteticModel`: Loads the aesthetic model. (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.
|
||||
- `AestheticScoreSorter`: Sorts the images by score.
|
||||
- `AestheticModel`: Loads the aesthetic model.
|
||||
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)
|
||||
- _Output_: String caption.
|
||||
- `FoodCategoryNode`: Classifies the food category of an image.
|
||||
- `FoodCategoryClassifierNode`: Classifies food categories in images.
|
||||
- _Input_: `image` (IMAGE)
|
||||
- _Output_: String category.
|
||||
- `AgeClassifierNode`: Classifies the age of a person in the image.
|
||||
- _Output_: Top 5 food categories with probabilities.
|
||||
- `AgeClassifierNode`: Classifies the age range in images.
|
||||
- _Input_: `image` (IMAGE)
|
||||
- _Output_: String age range.
|
||||
- `ImageClassifierNode`: General image classification.
|
||||
- _Output_: Top 5 age ranges with probabilities.
|
||||
- `ArtOrHumanClassifierNode`: Detects if an image is AI-generated or human-made.
|
||||
- _Input_: `image` (IMAGE), `show_on_node` (BOOL)
|
||||
- _Output_: String label, `artificial_prob` (INT), `human_prob` (INT)
|
||||
- `ClassifierNode`: A generic classifier node.
|
||||
- _Output_: Artificial and human probabilities.
|
||||
- `DocumentClassificationNode`: Classifies document types.
|
||||
- _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:
|
||||
- `SegformerNode`: Performs segmentation of the image.
|
||||
- _Input_: `image` (IMAGE), `model_name` (STRING), `show_on_node` (BOOL)
|
||||
- _Output_: Segmented image.
|
||||
- `SegformerNodeMasks`: Provides masks for the segmented images.
|
||||
- _Input_: No specific input detailed in the provided code.
|
||||
- _Output_: Image masks.
|
||||
- `SegformerNodeMergeSegments`: Merges certain segments in the segmented image.
|
||||
- _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING), `blur_radius` (INT), `dilation_radius` (INT), `intensity` (INT), `ceiling` (INT), `show_on_node` (BOOL)
|
||||
- _Output_: Image with merged segments.
|
||||
- `SeedIncrementerNode`: Increment the seed used for random processes.
|
||||
- _Input_: `seed` (INT), `increment_at` (INT)
|
||||
- _Output_: Incremented seed.
|
||||
- `StepCfgIncrementNode`: Calculates the step configuration for the process.
|
||||
- _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `img_steps` (INT), `max_steps` (INT)
|
||||
- _Output_: Calculated step configuration.
|
||||
3. **SegformerNode.py** - Handles semantic segmentation of images:
|
||||
- `SegformerNode`: Performs semantic segmentation with multiple model options.
|
||||
- _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, mask, info, and preview.
|
||||
- `SegformerNodeMasks`: Creates individual segment masks.
|
||||
- _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING)
|
||||
- _Output_: Image, mask, and segment info.
|
||||
- `SegformerNodeMergeSegments`: Merges and processes segments with advanced options.
|
||||
- _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_: Processed image, mask, info, and preview.
|
||||
- `SeedIncrementerNode`: Manages seed incrementation for workflows.
|
||||
- _Input_: `seed` (INT), `IncrementAt` (INT)
|
||||
- _Output_: Seed string, seed int, subseed string, subseed int.
|
||||
- `StepCfgIncrementNode`: Handles step and configuration increments.
|
||||
- _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `image_steps` (INT), `max_steps` (INT)
|
||||
- _Output_: CFG and steps values.
|
||||
|
||||
## Requirements
|
||||
|
||||
The project primarily uses the following libraries:
|
||||
The project requires the following Python libraries:
|
||||
|
||||
- Python
|
||||
- Torch
|
||||
- Transformers
|
||||
- PIL
|
||||
- Matplotlib
|
||||
- Numpy
|
||||
- IO
|
||||
- Scipy
|
||||
- torch
|
||||
- transformers
|
||||
- Pillow (PIL)
|
||||
- matplotlib
|
||||
- numpy
|
||||
- scipy
|
||||
- huggingface_hub
|
||||
- torchvision
|
||||
|
||||
## Installation
|
||||
|
||||
To install the necessary libraries, run:
|
||||
|
||||
1. Install the required Python packages:
|
||||
```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
|
||||
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,5 +1,6 @@
|
||||
#from .nodes.SegGPT import segGPTNode
|
||||
from .nodes import SegformerNode,ImageCaptioningNode,ImageProcessingNode
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
**SegformerNode.NODE_CLASS_MAPPINGS,
|
||||
|
||||
Binary file not shown.
@@ -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:
|
||||
@@ -127,7 +133,7 @@ class ArtOrHumanClassifierNode:
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"show_on_node": ("BOOL", {"default": False}),
|
||||
"show_on_node": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
OUTPUT_NODE = True
|
||||
@@ -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",
|
||||
}
|
||||
@@ -1,6 +1,5 @@
|
||||
import hashlib
|
||||
import fastapi
|
||||
import fastapi
|
||||
from fastapi import FastAPI
|
||||
import torch, time
|
||||
import io
|
||||
|
||||
@@ -241,60 +240,84 @@ class ImageFilterByFloatScoreNode:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"score": ("FLOAT", {"default": 0.0}),
|
||||
"threshold": ("FLOAT", {"default": 0.0}),
|
||||
"image": ("IMAGE", {"default": None}),
|
||||
"image": ("IMAGE",),
|
||||
"score": ("FLOAT", {"default": 0.0, "min": -100.0, "max": 100.0}),
|
||||
"threshold": ("FLOAT", {"default": 5.0, "min": -100.0, "max": 100.0}),
|
||||
"show_on_node": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "filter_image_by_score"
|
||||
CATEGORY = "LexTools/ImageProcessing/Scores"
|
||||
RETURN_TYPES = ("IMAGE", "FLOAT")
|
||||
FUNCTION = "filter_image"
|
||||
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:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"aesthetic_score": ("INT", {"default": None}),
|
||||
"ai_score_artificial": ("FLOAT", {"default": None}),
|
||||
"ai_score_human": ("FLOAT", {"default": None}),
|
||||
"show_on_node": ("INT", {"default": 0}),
|
||||
},
|
||||
"optional": {
|
||||
"image_score_good": ("FLOAT", {"default": 0}),
|
||||
"image_score_bad": ("FLOAT", {"default": 0}),
|
||||
"weight_good_score": ("FLOAT", {"default": 1}),
|
||||
"weight_aesthetic_score": ("FLOAT", {"default": 1.0}),
|
||||
"weight_bad_score": ("FLOAT", {"default": 1.0}),
|
||||
"weight_AIDetection": ("FLOAT", {"default": 1.0}),
|
||||
"weight_HumanDetection": ("FLOAT", {"default": 1.0}),
|
||||
"MultiplyScoreBy": ("FLOAT", {"default": 100000}),
|
||||
"aesthetic_score": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0}),
|
||||
"image_score_good": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0}),
|
||||
"image_score_bad": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.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}),
|
||||
"weight_good_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"weight_aesthetic_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"weight_bad_score": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"weight_AIDetection": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"weight_HumanDetection": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
|
||||
"MultiplyScoreBy": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0}),
|
||||
"show_on_node": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
OUTPUT_NODE = True
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
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):
|
||||
# Define the weights and maximum possible values
|
||||
maxA, maxB, maxC = 3, 3, 1000
|
||||
# Compute the exponential effect of the AI score
|
||||
ai_score_artificial_exp = 10 ** ai_score_artificial
|
||||
# Compute the final score according to the provided formula
|
||||
final_score = ((((((image_score_good + maxA) / (2 * maxA) * weight_good_score) + (aesthetic_score / maxC) * weight_bad_score) / (weight_good_score + weight_bad_score)) - weight_aesthetic_score * ((image_score_bad + maxB) / (2 * maxB))) * ((weight_HumanDetection * (ai_score_human))-( weight_AIDetection* (ai_score_artificial_exp)))) * MultiplyScoreBy
|
||||
def calculate_score(self, aesthetic_score, image_score_good, image_score_bad, ai_score_artificial, ai_score_human,
|
||||
weight_good_score, weight_aesthetic_score, weight_bad_score, weight_AIDetection, weight_HumanDetection,
|
||||
MultiplyScoreBy, show_on_node):
|
||||
try:
|
||||
# Calculate weighted scores
|
||||
weighted_aesthetic = float(aesthetic_score) * weight_aesthetic_score
|
||||
weighted_good = float(image_score_good) * weight_good_score
|
||||
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
|
||||
return (final_score, {"ui": {"STRING": [final_score]}})
|
||||
# Calculate total 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 {}}
|
||||
|
||||
|
||||
#
|
||||
@@ -376,7 +399,7 @@ class CalculateAestheticScore:
|
||||
"aesthetic_model": ("AESTHETIC_MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"keep_in_memory": ("BOOL", {"default": True}),
|
||||
"keep_in_memory": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -544,7 +567,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ScoreConverterNode":ScoreConverterNode,
|
||||
"MD5ImageHashNode": MD5ImageHashNode,
|
||||
"SamplerPropertiesNode": SamplerPropertiesNode,
|
||||
|
||||
"CalculateAestheticScore": CalculateAestheticScore,
|
||||
"LoadAesteticModel":AesteticModel,
|
||||
"AesthetlcScoreSorter": AesthetlcScoreSorter,
|
||||
}
|
||||
|
||||
|
||||
@@ -556,4 +581,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ScoreConverterNode":"Score Converter (Aesthetic Score)",
|
||||
"MD5ImageHashNode":"MD5 Image Hash",
|
||||
"SamplerPropertiesNode":"Sampler input node",
|
||||
"LoadAesteticModel": "LoadAesteticModel",
|
||||
"CalculateAestheticScore": "CalculateAestheticScore",
|
||||
"AesthetlcScoreSorter": "AesthetlcScoreSorter",
|
||||
}
|
||||
+435
-135
@@ -6,35 +6,99 @@ import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import io
|
||||
from scipy.ndimage import binary_dilation
|
||||
import os
|
||||
from pathlib import Path
|
||||
import json
|
||||
|
||||
|
||||
model_names = [
|
||||
"enes361/segformer_b2_clothes",
|
||||
"sayeed99/segformer_b3_clothes",
|
||||
"mattmdjaga/segformer_b0_clothes",
|
||||
"mattmdjaga/segformer_b2_clothes",
|
||||
"DiTo97/binarization-segformer-b3",
|
||||
"s3nh/SegFormer-b0-person-segmentation",
|
||||
"venture361/clothes_segmentation",
|
||||
"itsitgroup/human-body-segmentation",
|
||||
"matei-dorian/segformer-b5-finetuned-human-parsing",
|
||||
"Lexic0n/segformer-b0-finetuned-human-parsing",
|
||||
"sam1120/segformer-b0-finetuned-neurosymbolic-contingency-bag1-v0.1-v0",
|
||||
"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:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
global model_names # Assuming model_names is a list of model names
|
||||
return {
|
||||
"required": {
|
||||
"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"
|
||||
CATEGORY = "LexTools/ImageProcessing/Segmentation"
|
||||
|
||||
@@ -43,79 +107,199 @@ class SegformerNode:
|
||||
# self.processor = SegformerImageProcessor.from_pretrained("mattmdjaga/segformer_b2_clothes")
|
||||
# self.model = AutoModelForSemanticSegmentation.from_pretrained("mattmdjaga/segformer_b2_clothes")
|
||||
|
||||
def segment_image(self, image,model_name,):
|
||||
show_on_node = False
|
||||
self.processor = SegformerImageProcessor.from_pretrained(model_name)
|
||||
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)
|
||||
def process_mask(self, mask, normalize=True, binary=False, invert=False, post_process="none", radius=3):
|
||||
# Convert to float32 if not already
|
||||
mask = mask.float()
|
||||
|
||||
i = ImageOps.exif_transpose(img2)
|
||||
if i.getbands() != ("R", "G", "B", "A"):
|
||||
i = i.convert("RGBA")
|
||||
|
||||
# Normalize to 0-1 range if requested
|
||||
if normalize:
|
||||
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
|
||||
img2 = torch.from_numpy(img2)[None,]
|
||||
return mask
|
||||
|
||||
if 'A' in i.getbands():
|
||||
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
|
||||
# Get the unique segments in the image
|
||||
unique_segments = np.unique(pred_seg)
|
||||
def create_preview(self, image, mask):
|
||||
# Ensure image is in CHW format
|
||||
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()
|
||||
|
||||
# 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
|
||||
segment_info = []
|
||||
for segment in unique_segments:
|
||||
# Get the name of the segment from the model's configuration
|
||||
segment_name = self.model.config.id2label[segment]
|
||||
def parse_segment_groups(self, groups_str):
|
||||
if not groups_str.strip():
|
||||
return {}
|
||||
|
||||
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
|
||||
segment_info_str = "\n".join(segment_info)
|
||||
# Upsample logits with specified resize mode
|
||||
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:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
global model_names # Assuming model_names is a list of model names
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"default": None}),
|
||||
"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
|
||||
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
|
||||
show_on_node=False
|
||||
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
|
||||
i = 255. * image[0].cpu().numpy()
|
||||
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||
@@ -195,106 +382,219 @@ class SegformerNodeMasks:
|
||||
merged_image_pil = Image.fromarray(merged_image)
|
||||
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 {}
|
||||
|
||||
return {"result": (img2, merged_mask, 'Merged Segments'), "ui": output_ui}
|
||||
|
||||
|
||||
|
||||
class SegformerNodeMergeSegments:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
global model_names
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"default": None}),
|
||||
"segments_to_merge_str": ("STRING", {"default": ""}),
|
||||
"model_name": (model_names, {"default": model_names[0]}),
|
||||
"blur_radius": ("INT", {"default": 0}),
|
||||
"dilation_radius": ("INT", {"default": 0}), # Added dilation_radius
|
||||
"intensity": ("FLOAT", {"default": 1.0}), # Added intensity
|
||||
"ceiling": ("FLOAT", {"default": 1.0}), # Added ceiling
|
||||
|
||||
"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}),
|
||||
"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
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "STRING")
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "STRING", "IMAGE") # Added IMAGE for preview
|
||||
FUNCTION = "merge_segments"
|
||||
CATEGORY = "LexTools/ImageProcessing/Segmentation"
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def merge_segments(self, image, segments_to_merge_str, model_name, blur_radius, dilation_radius, intensity, ceiling): # Added dilation_radius in the arguments
|
||||
|
||||
show_on_node=False
|
||||
|
||||
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
|
||||
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:
|
||||
self.processor = SegformerImageProcessor.from_pretrained(model_name)
|
||||
except Exception:
|
||||
print(f"Failed to load preprocessor for model {model_name}. Using preprocessor from mattmdjaga/segformer_b2_clothes instead.")
|
||||
self.processor = SegformerImageProcessor.from_pretrained("matei-dorian/segformer-b5-finetuned-human-parsing")
|
||||
self.model = AutoModelForSemanticSegmentation.from_pretrained(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)
|
||||
|
||||
show_on_node = False
|
||||
|
||||
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].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))
|
||||
# Get input image dimensions and ensure proper shape
|
||||
input_image = image[0].cpu()
|
||||
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}")
|
||||
|
||||
img2 = np.array(merged_image_pil).astype(np.float32) / 255.0
|
||||
img2 = torch.from_numpy(img2).double()[None,]
|
||||
# Ensure image is in HWC format
|
||||
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 {}}
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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")
|
||||
@@ -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
@@ -1,4 +1,5 @@
|
||||
numpy
|
||||
opencv-python
|
||||
git+https://github.com/facebookresearch/detectron2.git
|
||||
pyodbc
|
||||
pyodbc
|
||||
pytorch_lightning
|
||||
|
||||
Reference in New Issue
Block a user