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.
This commit is contained in:
@@ -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.
|
||||
|
||||
|
||||
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:
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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 {}}
|
||||
|
||||
|
||||
#
|
||||
|
||||
+385
-115
@@ -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,27 +107,133 @@ 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,):
|
||||
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):
|
||||
# Create an RGBA preview with the mask as alpha channel
|
||||
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=""):
|
||||
# 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()
|
||||
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)
|
||||
logits = outputs.logits.cpu()
|
||||
|
||||
# Upsample logits with specified resize mode
|
||||
upsampled_logits = nn.functional.interpolate(
|
||||
logits,
|
||||
size=img.size[::-1],
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
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
|
||||
|
||||
# Convert the matplotlib figure to a PIL Image and return it
|
||||
# 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()
|
||||
plt.imshow(pred_seg)
|
||||
buf = io.BytesIO()
|
||||
@@ -71,51 +241,35 @@ class SegformerNode:
|
||||
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,]
|
||||
|
||||
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)
|
||||
# Create preview if requested
|
||||
preview = self.create_preview(image[0], merged_mask) if show_preview else None
|
||||
|
||||
# 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]
|
||||
|
||||
# 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
|
||||
# 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}
|
||||
|
||||
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 +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))
|
||||
@@ -195,106 +352,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 {}}
|
||||
|
||||
|
||||
|
||||
|
||||
+30
-6
@@ -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."
|
||||
version = "1.0.1"
|
||||
license = "LICENSE"
|
||||
dependencies = ["numpy", "opencv-python", "git+https://github.com/facebookresearch/detectron2.git", "pyodbc"]
|
||||
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"
|
||||
# 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