diff --git a/README.md b/README.md index 69b6c55..2ed7470 100644 --- a/README.md +++ b/README.md @@ -25,65 +25,92 @@ ComfyUI-LexTools is a Python-based image processing and analysis toolkit that us - _Output_: Converted score. Additional nodes from [GitHub Pages](https://github.com/strimmlarn/ComfyUI-Strimmlarns-Aesthetic-Score/) - These have been modified to improve performance and add an option to store the model in RAM, which significantly reduces generation time: - - `CalculateAestheticScore`: An optimized version of the original, with an option to keep the model loaded in RAM. (No specific input or output detailed in the provided code) - - `AesthetlcScoreSorter`: Sorts the images by score. (No specific input or output detailed in the provided code) - - `AesteticModel`: Loads the aesthetic model. (No specific input or output detailed in the provided code) + - `CalculateAestheticScore`: An optimized version of the original, with an option to keep the model loaded in RAM. + - `AesthetlcScoreSorter`: Sorts the images by score. + - `AesteticModel`: Loads the aesthetic model. 2. **ImageCaptioningNode.py** - Implements nodes for image captioning and classification: - - `ImageCaptioningNode`: Provides a caption for the image. + - `ImageCaptioningNode`: Provides a caption for the image using BLIP model. - _Input_: `image` (IMAGE) - _Output_: String caption. - - `FoodCategoryNode`: Classifies the food category of an image. + - `FoodCategoryClassifierNode`: Classifies food categories in images. - _Input_: `image` (IMAGE) - - _Output_: String category. - - `AgeClassifierNode`: Classifies the age of a person in the image. + - _Output_: Top 5 food categories with probabilities. + - `AgeClassifierNode`: Classifies the age range in images. - _Input_: `image` (IMAGE) - - _Output_: String age range. - - `ImageClassifierNode`: General image classification. + - _Output_: Top 5 age ranges with probabilities. + - `ArtOrHumanClassifierNode`: Detects if an image is AI-generated or human-made. - _Input_: `image` (IMAGE), `show_on_node` (BOOL) - - _Output_: String label, `artificial_prob` (INT), `human_prob` (INT) - - `ClassifierNode`: A generic classifier node. + - _Output_: Artificial and human probabilities. + - `DocumentClassificationNode`: Classifies document types. - _Input_: `image` (IMAGE) - - _Output_: String label. + - _Output_: Document type index and name. + - `NSFWClassifierNode`: Classifies content safety levels. + - _Input_: `image` (IMAGE), `show_on_node` (BOOL) + - _Output_: + - Classification report (STRING) + - NSFW Score (FLOAT) + - Neutral Score (FLOAT) + - Sexy Score (FLOAT) + - Porn Score (FLOAT) -3. **SegformerNode.py** - Handles semantic segmentation of images. It includes various nodes such as: - - `SegformerNode`: Performs segmentation of the image. - - _Input_: `image` (IMAGE), `model_name` (STRING), `show_on_node` (BOOL) - - _Output_: Segmented image. - - `SegformerNodeMasks`: Provides masks for the segmented images. - - _Input_: No specific input detailed in the provided code. - - _Output_: Image masks. - - `SegformerNodeMergeSegments`: Merges certain segments in the segmented image. - - _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING), `blur_radius` (INT), `dilation_radius` (INT), `intensity` (INT), `ceiling` (INT), `show_on_node` (BOOL) - - _Output_: Image with merged segments. - - `SeedIncrementerNode`: Increment the seed used for random processes. - - _Input_: `seed` (INT), `increment_at` (INT) - - _Output_: Incremented seed. - - `StepCfgIncrementNode`: Calculates the step configuration for the process. - - _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `img_steps` (INT), `max_steps` (INT) - - _Output_: Calculated step configuration. +3. **SegformerNode.py** - Handles semantic segmentation of images: + - `SegformerNode`: Performs semantic segmentation with multiple model options. + - _Input_: `image` (IMAGE), `model_name` (STRING), `normalize_mask` (BOOL), `binary_mask` (BOOL), `resize_mode` (STRING), `invert_mask` (BOOL), `show_preview` (BOOL), `return_individual_masks` (BOOL), `post_process` (STRING), `post_process_radius` (INT), `segment_groups` (STRING) + - _Output_: Segmented image, mask, info, and preview. + - `SegformerNodeMasks`: Creates individual segment masks. + - _Input_: `image` (IMAGE), `segments_to_merge` (STRING), `model_name` (STRING) + - _Output_: Image, mask, and segment info. + - `SegformerNodeMergeSegments`: Merges and processes segments with advanced options. + - _Input_: `image` (IMAGE), `segments_to_merge_str` (STRING), `model_name` (STRING), `normalize_mask` (BOOL), `binary_mask` (BOOL), `resize_mode` (STRING), `invert_mask` (BOOL), `show_preview` (BOOL), `blur_radius` (INT), `dilation_radius` (INT), `intensity` (FLOAT), `ceiling` (FLOAT) + - _Output_: Processed image, mask, info, and preview. + - `SeedIncrementerNode`: Manages seed incrementation for workflows. + - _Input_: `seed` (INT), `IncrementAt` (INT) + - _Output_: Seed string, seed int, subseed string, subseed int. + - `StepCfgIncrementNode`: Handles step and configuration increments. + - _Input_: `seed` (INT), `cfg_start` (INT), `steps_start` (INT), `image_steps` (INT), `max_steps` (INT) + - _Output_: CFG and steps values. ## Requirements -The project primarily uses the following libraries: +The project requires the following Python libraries: -- Python -- Torch -- Transformers -- PIL -- Matplotlib -- Numpy -- IO -- Scipy +- torch +- transformers +- Pillow (PIL) +- matplotlib +- numpy +- scipy +- huggingface_hub ## Installation -To install the necessary libraries, run: - +1. Install the required Python packages: ```bash -pip install torch transformers pillow matplotlib numpy scipy +pip install torch transformers pillow matplotlib numpy scipy huggingface_hub ``` +2. Clone this repository into your ComfyUI custom_nodes directory: +```bash +cd ComfyUI/custom_nodes +git clone https://github.com/YourUsername/ComfyUI-LexTools.git +``` + +3. Restart ComfyUI to load the new nodes. + +## Usage + +The nodes will appear in the ComfyUI interface under the "LexTools" category, organized into subcategories: +- LexTools/ImageProcessing/Segmentation +- LexTools/ImageProcessing/Classification +- LexTools/ImageProcessing/Captioning +- LexTools/Utilities + ## Contributing -Contributions to this project are welcome. If you find a bug or think of a feature that would benefit the project, please open an issue. If you'd like to contribute code, please open a pull request. + +Contributions are welcome! Please feel free to submit a Pull Request. For major changes, please open an issue first to discuss what you would like to change. + +## License + +This project is licensed under the MIT License - see the LICENSE file for details. diff --git a/models/watermark_model.pt b/models/watermark_model.pt new file mode 100644 index 0000000..d632507 Binary files /dev/null and b/models/watermark_model.pt differ diff --git a/nodes/ImageProcessingNode.py b/nodes/ImageProcessingNode.py index 8bd5dbc..d8279ff 100644 --- a/nodes/ImageProcessingNode.py +++ b/nodes/ImageProcessingNode.py @@ -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 {}} # diff --git a/nodes/SegformerNode.py b/nodes/SegformerNode.py index 4e4598a..df277a4 100644 --- a/nodes/SegformerNode.py +++ b/nodes/SegformerNode.py @@ -6,6 +6,9 @@ import matplotlib.pyplot as plt import numpy as np import io from scipy.ndimage import binary_dilation +import os +from pathlib import Path +import json model_names = [ @@ -15,20 +18,74 @@ model_names = [ "DiTo97/binarization-segformer-b3", "s3nh/SegFormer-b0-person-segmentation", "venture361/clothes_segmentation", + "itsitgroup/human-body-segmentation", "matei-dorian/segformer-b5-finetuned-human-parsing", "Lexic0n/segformer-b0-finetuned-human-parsing", "sam1120/segformer-b0-finetuned-neurosymbolic-contingency-bag1-v0.1-v0", "ehsanhallo/segformer-b0-scene-parse-150" ] +class SegformerModelLoader: + _models = {} # Cache for loaded models + _processors = {} # Cache for loaded processors + + @classmethod + def get_local_checkpoints(cls): + """Get list of local checkpoint directories""" + checkpoints_dir = Path("models/segformer") + if not checkpoints_dir.exists(): + checkpoints_dir.mkdir(parents=True, exist_ok=True) + return [] + + # Look for config.json files in subdirectories + checkpoints = [] + for path in checkpoints_dir.glob("*/config.json"): + checkpoints.append(path.parent.name) + return checkpoints + + @classmethod + def load_model(cls, model_name, local_dir=None): + """Load model and processor with caching""" + # Check cache first + cache_key = model_name if not local_dir else str(local_dir) + if cache_key in cls._models: + return cls._models[cache_key], cls._processors[cache_key] + + try: + if local_dir: + processor = SegformerImageProcessor.from_pretrained(local_dir) + model = AutoModelForSemanticSegmentation.from_pretrained(local_dir) + else: + processor = SegformerImageProcessor.from_pretrained(model_name) + model = AutoModelForSemanticSegmentation.from_pretrained(model_name) + + # Cache the loaded model and processor + cls._models[cache_key] = model + cls._processors[cache_key] = processor + return model, processor + except Exception as e: + print(f"Error loading model {model_name}: {str(e)}") + # Fallback to a reliable model + return cls.load_model("matei-dorian/segformer-b5-finetuned-human-parsing") + + @classmethod + def clear_cache(cls): + """Clear the model cache""" + cls._models.clear() + cls._processors.clear() + +# Update the model_names list to include local checkpoints +def get_available_models(): + local_checkpoints = SegformerModelLoader.get_local_checkpoints() + return model_names + [f"local:{cp}" for cp in local_checkpoints] + class SegformerNode: @classmethod def INPUT_TYPES(cls): - global model_names return { "required": { "image": ("IMAGE", {"default": None}), - "model_name": (model_names, {"default": model_names[0]}), + "model_name": (get_available_models(), {"default": model_names[0]}), "normalize_mask": ("BOOLEAN", {"default": True}), "binary_mask": ("BOOLEAN", {"default": False}), "resize_mode": (["nearest", "bilinear", "bicubic"], {"default": "bilinear"}), @@ -113,9 +170,14 @@ class SegformerNode: resize_mode="bilinear", invert_mask=False, show_preview=True, return_individual_masks=False, post_process="none", post_process_radius=3, segment_groups=""): + # Handle local checkpoint loading + if model_name.startswith("local:"): + local_dir = Path("models/segformer") / model_name[6:] + self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir) + else: + self.model, self.processor = SegformerModelLoader.load_model(model_name) + show_on_node = False - self.processor = SegformerImageProcessor.from_pretrained(model_name) - self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name) # Process input image i = 255. * image[0].cpu().numpy() @@ -203,13 +265,11 @@ class SegformerNode: class SegformerNodeMasks: @classmethod def INPUT_TYPES(cls): - global model_names # Assuming model_names is a list of model names return { "required": { "image": ("IMAGE", {"default": None}), "segments_to_merge": ("STRING", {"default": "0"}), - "model_name": (model_names, {"default": model_names[0]}) - + "model_name": (get_available_models(), {"default": model_names[0]}), }, } @@ -222,14 +282,17 @@ class SegformerNodeMasks: # Function to segment the image and return the merged segments as per the provided indices def segment_image(self, image, segments_to_merge, model_name): + # Handle local checkpoint loading + if model_name.startswith("local:"): + local_dir = Path("models/segformer") / model_name[6:] + self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir) + else: + self.model, self.processor = SegformerModelLoader.load_model(model_name) + # Convert the segments_to_merge from string to list of integers show_on_node=False segments_to_merge = list(map(int, segments_to_merge.split(','))) - # Load the pretrained models and processors - self.processor = SegformerImageProcessor.from_pretrained(model_name) - self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name) - # Preprocess the image i = 255. * image[0].cpu().numpy() img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) @@ -300,12 +363,11 @@ class SegformerNodeMasks: class SegformerNodeMergeSegments: @classmethod def INPUT_TYPES(cls): - global model_names return { "required": { "image": ("IMAGE", {"default": None}), "segments_to_merge_str": ("STRING", {"default": ""}), - "model_name": (model_names, {"default": model_names[0]}), + "model_name": (get_available_models(), {"default": model_names[0]}), "normalize_mask": ("BOOLEAN", {"default": True}), "binary_mask": ("BOOLEAN", {"default": False}), "resize_mode": (["nearest", "bilinear", "bicubic"], {"default": "bilinear"}), @@ -332,6 +394,10 @@ class SegformerNodeMergeSegments: mask = torch.from_numpy(mask) mask = mask.float() + # Ensure mask is 2D + if len(mask.shape) > 2: + mask = mask.squeeze() + # Normalize to 0-1 range if requested if normalize: min_val = mask.min() @@ -355,8 +421,9 @@ class SegformerNodeMergeSegments: # Apply Gaussian blur for feathering if blur_radius > 0: - mask_np = (mask.numpy() * 255).astype('uint8') - mask_pil = Image.fromarray(mask_np) + # Ensure mask is 2D and in correct range for PIL + mask_np = (mask.squeeze().numpy() * 255).astype(np.uint8) + mask_pil = Image.fromarray(mask_np, mode='L') # Use 'L' mode for grayscale mask_pil = mask_pil.filter(ImageFilter.GaussianBlur(radius=blur_radius)) mask = torch.from_numpy(np.array(mask_pil).astype(np.float32) / 255.0) @@ -373,6 +440,24 @@ class SegformerNodeMergeSegments: # Create an RGBA preview with the mask as alpha channel if len(image.shape) == 2: image = image.unsqueeze(0).repeat(3, 1, 1) + elif len(image.shape) == 3: + if image.shape[0] != 3: # If channels are not in first dimension + image = image.permute(2, 0, 1) # Move channels to first dimension + + # Ensure mask has correct dimensions + if len(mask.shape) == 3: + mask = mask.squeeze(0) + if len(mask.shape) > 2: + mask = mask.squeeze() + + if mask.shape != image.shape[1:]: + mask = torch.nn.functional.interpolate( + mask.unsqueeze(0).unsqueeze(0), + size=image.shape[1:], + mode='bilinear', + align_corners=False + ).squeeze() + preview = image.clone() preview = torch.cat([preview, mask.unsqueeze(0)], dim=0) return preview @@ -381,76 +466,105 @@ class SegformerNodeMergeSegments: binary_mask=False, resize_mode="bilinear", invert_mask=False, show_preview=True, blur_radius=5, dilation_radius=5, intensity=1.0, ceiling=1.0): - show_on_node = False - try: - self.processor = SegformerImageProcessor.from_pretrained(model_name) - except Exception: - print(f"Failed to load preprocessor for model {model_name}. Using preprocessor from mattmdjaga/segformer_b2_clothes instead.") - self.processor = SegformerImageProcessor.from_pretrained("matei-dorian/segformer-b5-finetuned-human-parsing") - self.model = AutoModelForSemanticSegmentation.from_pretrained(model_name) + # Handle local checkpoint loading + if model_name.startswith("local:"): + local_dir = Path("models/segformer") / model_name[6:] + self.model, self.processor = SegformerModelLoader.load_model(model_name, local_dir) + else: + self.model, self.processor = SegformerModelLoader.load_model(model_name) + + show_on_node = False - i = 255. * image[0].cpu().numpy() - img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) - inputs = self.processor(images=img, return_tensors="pt") + # Get input image dimensions and ensure proper shape + input_image = image[0].cpu() + if len(input_image.shape) != 3: + raise ValueError(f"Expected input image with shape (H,W,C) or (C,H,W), got {input_image.shape}") + + # Ensure image is in HWC format + if input_image.shape[0] == 3: # If in CHW format + input_image = input_image.permute(1, 2, 0) + + input_height, input_width = input_image.shape[0:2] + + # Process input image + img = Image.fromarray((input_image.numpy() * 255).astype(np.uint8)) + inputs = self.processor(images=img, return_tensors="pt") - outputs = self.model(**inputs) - logits = outputs.logits.cpu() + outputs = self.model(**inputs) + logits = outputs.logits.cpu() - upsampled_logits = nn.functional.interpolate( - logits, - size=img.size[::-1], - mode=resize_mode, - align_corners=False if resize_mode != "nearest" else None, - ) + # Upsample logits to match input image size + upsampled_logits = nn.functional.interpolate( + logits, + size=(input_height, input_width), + mode=resize_mode, + align_corners=False if resize_mode != "nearest" else None, + ) - pred_seg = upsampled_logits.argmax(dim=1)[0].numpy() - unique_segments = np.unique(pred_seg) + pred_seg = upsampled_logits.argmax(dim=1)[0].numpy() + unique_segments = np.unique(pred_seg) - # Handle empty segments string - if not segments_to_merge_str.strip(): - segments_to_merge = [] - else: - segments_to_merge = [int(s.strip()) for s in segments_to_merge_str.split(',') if s.strip()] + # Handle empty segments string + if not segments_to_merge_str.strip(): + segments_to_merge = [] + else: + segments_to_merge = [int(s.strip()) for s in segments_to_merge_str.split(',') if s.strip()] - merged_mask = np.zeros_like(pred_seg, dtype=np.float32) - merged_segments = [] + merged_mask = np.zeros((input_height, input_width), dtype=np.float32) + merged_segments = [] - for segment in unique_segments: - if segment in segments_to_merge: - mask = np.where(pred_seg == segment, 1, 0) - merged_mask = np.maximum(merged_mask, mask) - merged_segments.append(segment) + for segment in unique_segments: + if segment in segments_to_merge: + mask = np.where(pred_seg == segment, 1, 0) + merged_mask = np.maximum(merged_mask, mask) + merged_segments.append(segment) - # Convert to tensor and process - merged_mask = torch.from_numpy(merged_mask) - merged_mask = self.process_mask( - merged_mask, - normalize=normalize_mask, - binary=binary_mask, - invert=invert_mask, - blur_radius=blur_radius, - dilation_radius=dilation_radius, - intensity=intensity, - ceiling=ceiling - ) + # Convert to tensor and process + merged_mask = torch.from_numpy(merged_mask) + merged_mask = self.process_mask( + merged_mask, + normalize=normalize_mask, + binary=binary_mask, + invert=invert_mask, + blur_radius=blur_radius, + dilation_radius=dilation_radius, + intensity=intensity, + ceiling=ceiling + ) - # Create preview if requested - preview = self.create_preview(image[0], merged_mask) if show_preview else None + # Ensure mask has correct dimensions for broadcasting + merged_mask_3d = merged_mask.unsqueeze(-1) # Add channel dimension for broadcasting - # Apply mask to image - merged_image = image[0].cpu().numpy() * merged_mask.numpy()[..., None] - merged_image = torch.from_numpy(merged_image).unsqueeze(0) + # Apply mask to image + merged_image = input_image.numpy() * merged_mask_3d.numpy() + + # Convert back to tensor in CHW format + merged_image = torch.from_numpy(merged_image).permute(2, 0, 1).unsqueeze(0) - merged_segments_str = ','.join(map(str, merged_segments)) - if not merged_segments: - merged_segments_str = "No segments selected" + merged_segments_str = ','.join(map(str, merged_segments)) + if not merged_segments: + merged_segments_str = "No segments selected" - output_ui = {"images": [merged_image]} if show_on_node else {} + # Create preview + if show_preview: + preview = self.create_preview(input_image.permute(2, 0, 1), merged_mask) + else: + preview = merged_image - return {"result": (merged_image, merged_mask, merged_segments_str, - preview if preview is not None else merged_image), - "ui": output_ui} + output_ui = {"images": [merged_image]} if show_on_node else {} + + return {"result": (merged_image, merged_mask, merged_segments_str, preview), + "ui": output_ui} + + except Exception as e: + import traceback + print(f"Error merging segments: {str(e)}") + print(f"Traceback: {traceback.format_exc()}") + # Return original image and empty mask on error + empty_mask = torch.zeros((input_height, input_width), dtype=torch.float32) + return {"result": (image, empty_mask, f"Error: {str(e)}", image), + "ui": {"images": [image]} if show_on_node else {}} diff --git a/pyproject.toml b/pyproject.toml index f04c9ce..3e774e9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,15 +1,39 @@ [project] name = "comfyui-lextools" -description = "ComfyUI-LexTools is a Python-based image processing and analysis toolkit that uses machine learning models for semantic image segmentation, image scoring, and image captioning." +description = """ +A comprehensive toolkit for ComfyUI that provides advanced image processing, analysis, and AI-powered features: +- Semantic segmentation with multiple pre-trained models and mask processing +- Image classification (age, food, documents, NSFW content, AI detection) +- Image captioning using BLIP +- Image quality scoring and filtering +- Workflow utilities for seed management and image aspect ratio handling +""" version = "1.0.1" -license = "LICENSE" -dependencies = ["numpy", "opencv-python", "git+https://github.com/facebookresearch/detectron2.git", "pyodbc"] +license = "MIT" +dependencies = [ + "torch>=2.0.0", + "transformers>=4.30.0", + "Pillow>=9.0.0", + "matplotlib>=3.0.0", + "numpy>=1.20.0", + "scipy>=1.7.0", + "huggingface_hub>=0.19.0" +] [project.urls] Repository = "https://github.com/SOELexicon/ComfyUI-LexTools" -# Used by Comfy Registry https://comfyregistry.org +Documentation = "https://github.com/SOELexicon/ComfyUI-LexTools/blob/main/README.md" [tool.comfy] PublisherId = "lexicon" DisplayName = "ComfyUI-LexTools" -Icon = "" +Description = "Advanced image processing and AI analysis toolkit for ComfyUI" +Icon = "🛠️" +Tags = [ + "image processing", + "segmentation", + "classification", + "captioning", + "workflow", + "utilities" +]