Author SHA1 Message Date
Craig Wright be62884389 Merge pull request #15 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-03-28 10:50:35 +00:00
Craig Wright e932dccb88 Merge pull request #14 from SOELexicon/segformer-updates
fix: Resolve tensor dimension mismatch in SegformerNode - Improve mas…
2025-03-26 22:00:49 +00:00
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 Wright 44288b9cb4 Merge pull request #13 from SOELexicon/segformer-updates
This PR updates various image processing nodes and improves configuration details for ComfyUI-LexTools. Key changes include:

Modifications to the ImageProcessingNode and its filtering function to support a wider score range and additional UI output.
Enhancements to the image captioning and classification nodes, including return type adjustments and new NSFW and watermark detection nodes.
Updates to the project configuration and README to reflect new features and dependency versions.
2025-03-26 17:46:33 +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
snomiao c39ca74028 chore(publish): update GitHub Actions workflow for node publishing
- Add permissions for writing issues
- Update action version to v1 for publish-node-action
- Add condition to run job only for 'SOELexicon' repository owner
2025-01-21 08:46:42 +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 923 additions and 225 deletions
+25
View File
@@ -0,0 +1,25 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'SOELexicon' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
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",
} }
+68 -40
View File
@@ -1,6 +1,5 @@
import hashlib import hashlib
import fastapi from fastapi import FastAPI
import fastapi
import torch, time import torch, time
import io import io
@@ -241,60 +240,84 @@ class ImageFilterByFloatScoreNode:
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): def filter_image(self, image, score, threshold, show_on_node):
# If score > threshold, return the image, otherwise return None try:
if score < threshold: if float(score) >= float(threshold):
pass 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: else:
return (image,) 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)) - weight_aesthetic_score * ((image_score_bad + maxB) / (2 * maxB))) * ((weight_HumanDetection * (ai_score_human))-( weight_AIDetection* (ai_score_artificial_exp)))) * MultiplyScoreBy 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 # 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 {}}
# #
@@ -376,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}),
} }
} }
@@ -544,7 +567,9 @@ NODE_CLASS_MAPPINGS = {
"ScoreConverterNode":ScoreConverterNode, "ScoreConverterNode":ScoreConverterNode,
"MD5ImageHashNode": MD5ImageHashNode, "MD5ImageHashNode": MD5ImageHashNode,
"SamplerPropertiesNode": SamplerPropertiesNode, "SamplerPropertiesNode": SamplerPropertiesNode,
"CalculateAestheticScore": CalculateAestheticScore,
"LoadAesteticModel":AesteticModel,
"AesthetlcScoreSorter": AesthetlcScoreSorter,
} }
@@ -556,4 +581,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ScoreConverterNode":"Score Converter (Aesthetic Score)", "ScoreConverterNode":"Score Converter (Aesthetic Score)",
"MD5ImageHashNode":"MD5 Image Hash", "MD5ImageHashNode":"MD5 Image Hash",
"SamplerPropertiesNode":"Sampler input node", "SamplerPropertiesNode":"Sampler input node",
"LoadAesteticModel": "LoadAesteticModel",
"CalculateAestheticScore": "CalculateAestheticScore",
"AesthetlcScoreSorter": "AesthetlcScoreSorter",
} }
+389 -89
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,27 +107,156 @@ 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):
# Convert to float32 if not already
mask = mask.float()
# 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
return mask
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
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
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 show_on_node = False
self.processor = SegformerImageProcessor.from_pretrained(model_name)
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name) # Process input 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))
inputs = self.processor(images=img, return_tensors="pt") inputs = self.processor(images=img, return_tensors="pt")
# Get model outputs
outputs = self.model(**inputs) outputs = self.model(**inputs)
logits = outputs.logits.cpu() logits = outputs.logits.cpu()
# Upsample logits with specified resize mode
upsampled_logits = nn.functional.interpolate( upsampled_logits = nn.functional.interpolate(
logits, logits,
size=img.size[::-1], size=img.size[::-1],
mode="bilinear", mode=resize_mode,
align_corners=False, align_corners=False if resize_mode != "nearest" else None,
) )
pred_seg = upsampled_logits.argmax(dim=1)[0] pred_seg = upsampled_logits.argmax(dim=1)[0]
# Convert the matplotlib figure to a PIL Image and return it # 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
# 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)
# 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() fig = plt.figure()
plt.imshow(pred_seg) plt.imshow(pred_seg)
buf = io.BytesIO() buf = io.BytesIO()
@@ -71,51 +264,42 @@ class SegformerNode:
buf.seek(0) buf.seek(0)
img2 = Image.open(buf) img2 = Image.open(buf)
# Convert visualization to tensor
i = ImageOps.exif_transpose(img2) i = ImageOps.exif_transpose(img2)
if i.getbands() != ("R", "G", "B", "A"): if i.getbands() != ("R", "G", "B", "A"):
i = i.convert("RGBA") i = i.convert("RGBA")
img2 = np.array(img2).astype(np.float32) / 255.0 img2 = np.array(img2).astype(np.float32) / 255.0
img2 = torch.from_numpy(img2)[None,] img2 = torch.from_numpy(img2)[None,]
if 'A' in i.getbands(): # Create preview if requested
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 preview = self.create_preview(image[0], merged_mask) if show_preview else None
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)
# Create a string with the information for each segment # Join segment info
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]
# Here, you would replace these values with the actual accuracy and IoU for the segment
segment_info.append(f"Segment {segment}: {segment_name}")
# Join the segment info strings into a single string
segment_info_str = "\n".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 {} output_ui = {"images": [img2]} if show_on_node else {}
return {"result": (img2,mask, segment_info_str), "ui": output_ui} # 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
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:
# 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 show_on_node = False
try: # Get input image dimensions and ensure proper shape
self.processor = SegformerImageProcessor.from_pretrained(model_name) input_image = image[0].cpu()
except Exception: if len(input_image.shape) != 3:
print(f"Failed to load preprocessor for model {model_name}. Using preprocessor from mattmdjaga/segformer_b2_clothes instead.") raise ValueError(f"Expected input image with shape (H,W,C) or (C,H,W), got {input_image.shape}")
self.processor = SegformerImageProcessor.from_pretrained("matei-dorian/segformer-b5-finetuned-human-parsing")
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name)
i = 255. * image[0].cpu().numpy() # Ensure image is in HWC format
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) 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") inputs = self.processor(images=img, return_tensors="pt")
outputs = self.model(**inputs) outputs = self.model(**inputs)
logits = outputs.logits.cpu() logits = outputs.logits.cpu()
# Upsample logits to match input image size
upsampled_logits = nn.functional.interpolate( upsampled_logits = nn.functional.interpolate(
logits, logits,
size=img.size[::-1], size=(input_height, input_width),
mode="bilinear", mode=resize_mode,
align_corners=False, align_corners=False if resize_mode != "nearest" else None,
) )
pred_seg = upsampled_logits.argmax(dim=1)[0].numpy() pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
unique_segments = np.unique(pred_seg) unique_segments = np.unique(pred_seg)
segments_to_merge = list(map(int, segments_to_merge_str.split(','))) # Handle empty segments string
if not segments_to_merge_str.strip():
merged_mask = np.zeros_like(pred_seg) 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 = [] merged_segments = []
for segment in unique_segments: for segment in unique_segments:
if segment in segments_to_merge: if segment in segments_to_merge:
mask = np.where(pred_seg == segment, 1, 0) 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_mask = np.maximum(merged_mask, mask)
merged_segments.append(segment) merged_segments.append(segment)
merged_mask = np.clip(merged_mask * intensity, 0, ceiling) # Apply intensity and ceiling to the mask # Convert to tensor and process
if dilation_radius > 0: # Dilate the mask if dilation_radius > 0 merged_mask = torch.from_numpy(merged_mask)
struct = np.ones((2 * dilation_radius + 1, 2 * dilation_radius + 1)) merged_mask = self.process_mask(
merged_mask = binary_dilation(merged_mask, structure=struct) merged_mask,
merged_mask_rgb = np.repeat(merged_mask[..., None], 3, axis=2) normalize=normalize_mask,
if blur_radius > 0: # Blur the mask if radius > 0 binary=binary_mask,
merged_mask_rgb = Image.fromarray((merged_mask_rgb * 255).astype('uint8')) invert=invert_mask,
merged_mask_rgb = merged_mask_rgb.filter(ImageFilter.GaussianBlur(radius=blur_radius)) blur_radius=blur_radius,
merged_mask_rgb = np.array(merged_mask_rgb) / 255.0 dilation_radius=dilation_radius,
intensity=intensity,
ceiling=ceiling
)
merged_image = np.array(img) * merged_mask_rgb # Ensure mask has correct dimensions for broadcasting
merged_mask_3d = merged_mask.unsqueeze(-1) # Add channel dimension for broadcasting
merged_image_pil = Image.fromarray(merged_image.astype('uint8')) # Apply mask to image
if blur_radius > 0: # Apply blur if radius > 0 merged_image = input_image.numpy() * merged_mask_3d.numpy()
merged_image_pil = merged_image_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius))
img2 = np.array(merged_image_pil).astype(np.float32) / 255.0 # Convert back to tensor in CHW format
img2 = torch.from_numpy(img2).double()[None,] merged_image = torch.from_numpy(merged_image).permute(2, 0, 1).unsqueeze(0)
merged_mask_torch = torch.from_numpy(merged_mask).float()[None,] # change from double to float
merged_segments_str = ','.join(map(str, merged_segments)) merged_segments_str = ','.join(map(str, merged_segments))
if not merged_segments:
merged_segments_str = "No segments selected"
output_ui = {"images": [img2]} if show_on_node else {} # Create preview
if show_preview:
preview = self.create_preview(input_image.permute(2, 0, 1), merged_mask)
else:
preview = merged_image
return {"result": (img2, merged_mask_torch, merged_segments_str), "ui": output_ui} 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"
]
+1
View File
@@ -2,3 +2,4 @@ numpy
opencv-python opencv-python
git+https://github.com/facebookresearch/detectron2.git git+https://github.com/facebookresearch/detectron2.git
pyodbc pyodbc
pytorch_lightning