feat: Update WatermarkDetectionNode with improved model handling and error recovery
This commit is contained in:
@@ -25,65 +25,92 @@ 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.
|
||||
- `AesthetlcScoreSorter`: Sorts the images by score.
|
||||
- `AesteticModel`: 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)
|
||||
- _Output_:
|
||||
- Classification report (STRING)
|
||||
- NSFW Score (FLOAT)
|
||||
- Neutral Score (FLOAT)
|
||||
- Sexy Score (FLOAT)
|
||||
- Porn Score (FLOAT)
|
||||
|
||||
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
|
||||
|
||||
## 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
|
||||
```
|
||||
|
||||
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.
|
||||
|
||||
|
||||
Binary file not shown.
@@ -240,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 {}}
|
||||
|
||||
|
||||
#
|
||||
|
||||
+185
-71
@@ -6,6 +6,9 @@ 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 = [
|
||||
@@ -15,20 +18,74 @@ model_names = [
|
||||
"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
|
||||
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"}),
|
||||
@@ -113,9 +170,14 @@ class SegformerNode:
|
||||
resize_mode="bilinear", invert_mask=False, show_preview=True,
|
||||
return_individual_masks=False, post_process="none",
|
||||
post_process_radius=3, segment_groups=""):
|
||||
# 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
|
||||
self.processor = SegformerImageProcessor.from_pretrained(model_name)
|
||||
self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name)
|
||||
|
||||
# Process input image
|
||||
i = 255. * image[0].cpu().numpy()
|
||||
@@ -203,13 +265,11 @@ class SegformerNode:
|
||||
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]}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -222,14 +282,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))
|
||||
@@ -300,12 +363,11 @@ class SegformerNodeMasks:
|
||||
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]}),
|
||||
"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"}),
|
||||
@@ -332,6 +394,10 @@ class SegformerNodeMergeSegments:
|
||||
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()
|
||||
@@ -355,8 +421,9 @@ class SegformerNodeMergeSegments:
|
||||
|
||||
# Apply Gaussian blur for feathering
|
||||
if blur_radius > 0:
|
||||
mask_np = (mask.numpy() * 255).astype('uint8')
|
||||
mask_pil = Image.fromarray(mask_np)
|
||||
# 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)
|
||||
|
||||
@@ -373,6 +440,24 @@ class SegformerNodeMergeSegments:
|
||||
# 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
|
||||
@@ -381,76 +466,105 @@ class SegformerNodeMergeSegments:
|
||||
binary_mask=False, resize_mode="bilinear", invert_mask=False,
|
||||
show_preview=True, blur_radius=5, dilation_radius=5,
|
||||
intensity=1.0, ceiling=1.0):
|
||||
show_on_node = False
|
||||
|
||||
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")
|
||||
# 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}")
|
||||
|
||||
# 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")
|
||||
|
||||
outputs = self.model(**inputs)
|
||||
logits = outputs.logits.cpu()
|
||||
outputs = self.model(**inputs)
|
||||
logits = outputs.logits.cpu()
|
||||
|
||||
upsampled_logits = nn.functional.interpolate(
|
||||
logits,
|
||||
size=img.size[::-1],
|
||||
mode=resize_mode,
|
||||
align_corners=False if resize_mode != "nearest" else None,
|
||||
)
|
||||
# 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,
|
||||
)
|
||||
|
||||
pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
|
||||
unique_segments = np.unique(pred_seg)
|
||||
pred_seg = upsampled_logits.argmax(dim=1)[0].numpy()
|
||||
unique_segments = np.unique(pred_seg)
|
||||
|
||||
# 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()]
|
||||
# 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_like(pred_seg, dtype=np.float32)
|
||||
merged_segments = []
|
||||
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)
|
||||
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
|
||||
)
|
||||
# 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
|
||||
)
|
||||
|
||||
# Create preview if requested
|
||||
preview = self.create_preview(image[0], merged_mask) if show_preview else None
|
||||
# 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 = image[0].cpu().numpy() * merged_mask.numpy()[..., None]
|
||||
merged_image = torch.from_numpy(merged_image).unsqueeze(0)
|
||||
# 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"
|
||||
merged_segments_str = ','.join(map(str, merged_segments))
|
||||
if not merged_segments:
|
||||
merged_segments_str = "No segments selected"
|
||||
|
||||
output_ui = {"images": [merged_image]} 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": (merged_image, merged_mask, merged_segments_str,
|
||||
preview if preview is not None else merged_image),
|
||||
"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 {}}
|
||||
|
||||
|
||||
|
||||
|
||||
+29
-5
@@ -1,15 +1,39 @@
|
||||
[project]
|
||||
name = "comfyui-lextools"
|
||||
description = "ComfyUI-LexTools is a Python-based image processing and analysis toolkit that uses machine learning models for semantic image segmentation, image scoring, and image captioning."
|
||||
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.1"
|
||||
license = "LICENSE"
|
||||
dependencies = ["numpy", "opencv-python", "git+https://github.com/facebookresearch/detectron2.git", "pyodbc"]
|
||||
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"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
Documentation = "https://github.com/SOELexicon/ComfyUI-LexTools/blob/main/README.md"
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "lexicon"
|
||||
DisplayName = "ComfyUI-LexTools"
|
||||
Icon = ""
|
||||
Description = "Advanced image processing and AI analysis toolkit for ComfyUI"
|
||||
Icon = "🛠️"
|
||||
Tags = [
|
||||
"image processing",
|
||||
"segmentation",
|
||||
"classification",
|
||||
"captioning",
|
||||
"workflow",
|
||||
"utilities"
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user