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