diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..c199ffb --- /dev/null +++ b/__init__.py @@ -0,0 +1,8 @@ +#from .nodes.SegGPT import segGPTNode +from .nodes import SegformerNode,ImageCaptioningNode,ImageProcessingNode +NODE_CLASS_MAPPINGS = { + + **SegformerNode.NODE_CLASS_MAPPINGS, + **ImageCaptioningNode.NODE_CLASS_MAPPINGS, + **ImageProcessingNode.NODE_CLASS_MAPPINGS, +} \ No newline at end of file diff --git a/nodes/ImageCaptioningNode.py b/nodes/ImageCaptioningNode.py new file mode 100644 index 0000000..726d825 --- /dev/null +++ b/nodes/ImageCaptioningNode.py @@ -0,0 +1,209 @@ +import time +import torch +from transformers import BlipProcessor,AutoModel, BlipForConditionalGeneration,AutoFeatureExtractor,AutoModelForImageClassification,ViTFeatureExtractor, ViTForImageClassification, AutoModelForImageClassification +from PIL import Image +import numpy as np +from scipy.ndimage import binary_dilation + + +class ImageCaptioningNode: + @classmethod + def INPUT_TYPES(s): + return {"required": {"image": ("IMAGE",)}} + + RETURN_TYPES = ("STRING",) + FUNCTION = "caption" + + CATEGORY = "LexTools/ImageProcessing/Captioning" + + def __init__(self): + self.processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-large") + self.model = BlipForConditionalGeneration.from_pretrained("Salesforce/blip-image-captioning-large").to("cuda") + + def caption(self, image): + image = image.numpy() + if image.ndim == 4: # image has batch dimension + image = image[0] # take first image in batch + image = Image.fromarray((image * 255).astype(np.uint8).transpose(1, 2, 0)) + + # Perform unconditional image captioning + inputs = self.processor(image, return_tensors="pt").to("cuda") + out = self.model.generate(**inputs) + caption = self.processor.decode(out[0], skip_special_tokens=True) + + return (caption,) + + +class FoodCategoryClassifierNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", {"default": None}), + }, + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "classify_FoodCategory" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def __init__(self): + self.feature_extractor = AutoFeatureExtractor.from_pretrained('Kaludi/food-category-classification-v2.0') + self.model = AutoModelForImageClassification.from_pretrained('Kaludi/food-category-classification-v2.0') + + def classify_FoodCategory(self, image): + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + inputs = self.feature_extractor(images=img, return_tensors="pt") + + outputs = self.model(**inputs) + proba = outputs.logits.softmax(1) + + # Get the top 5 class probabilities and their indices + top_5_probs, top_5_indices = torch.topk(proba, 5) + + # Convert the probabilities and indices to lists + top_5_probs = top_5_probs.tolist()[0] + top_5_indices = top_5_indices.tolist()[0] + + # Get the labels from the model's configuration + labels = self.model.config.id2label + + # Create a list of dictionaries with the class labels and probabilities + results = [{"score": prob, "label": labels[idx]} for prob, idx in zip(top_5_probs, top_5_indices)] + + # Convert the list of dictionaries to a string and return it + results_str = str(results) + return [results_str] + +class AgeClassifierNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", {"default": None}), + }, + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "classify_age" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def __init__(self): + self.feature_extractor = ViTFeatureExtractor.from_pretrained('nateraw/vit-age-classifier') + self.model = ViTForImageClassification.from_pretrained('nateraw/vit-age-classifier') + + def classify_age(self, image): + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + inputs = self.feature_extractor(images=img, return_tensors="pt") + + outputs = self.model(**inputs) + proba = outputs.logits.softmax(1) + + # Get the top 5 class probabilities and their indices + top_5_probs, top_5_indices = torch.topk(proba, 5) + + # Convert the probabilities and indices to lists + top_5_probs = top_5_probs.tolist()[0] + top_5_indices = top_5_indices.tolist()[0] + + # Get the labels from the model's configuration + labels = self.model.config.id2label + + # Create a list of dictionaries with the class labels and probabilities + results = [{"score": prob, "label": labels[idx]} for prob, idx in zip(top_5_probs, top_5_indices)] + + # Convert the list of dictionaries to a string and return it + results_str = str(results) + return [results_str] + + +class ArtOrHumanClassifierNode: + model = None + feature_extractor = None + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "show_on_node": ("BOOL", {"default": False}), + }, + } + OUTPUT_NODE = True + + RETURN_TYPES = ("FLOAT", "FLOAT") + FUNCTION = "classify_image" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def __init__(self): + if not ArtOrHumanClassifierNode.feature_extractor: + ArtOrHumanClassifierNode.feature_extractor = ViTFeatureExtractor.from_pretrained("umm-maybe/AI-image-detector") + + if not ArtOrHumanClassifierNode.model: + ArtOrHumanClassifierNode.model = AutoModelForImageClassification.from_pretrained("umm-maybe/AI-image-detector") + + def classify_image(self, image, show_on_node): + i = 255. * image[0].cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + inputs = self.feature_extractor(images=img, return_tensors="pt") + + outputs = self.model(**inputs) + 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() + + output_ui = {"text": [artificial_prob]} if show_on_node else {} + + return {"result": (artificial_prob, human_prob), "ui": output_ui} + + + + +class DocumentClassificationNode: + @classmethod + def INPUT_TYPES(cls): + return {"required": {"image": ("IMAGE", {"default": None})}} + + RETURN_TYPES = ("INT", "STRING") + FUNCTION = "classify" + CATEGORY = "LexTools/ImageProcessing/Classification" + + def __init__(self): + self.extractor = AutoFeatureExtractor.from_pretrained("DunnBC22/dit-base-Document_Classification-RVL_CDIP") + self.model = AutoModelForImageClassification.from_pretrained("DunnBC22/dit-base-Document_Classification-RVL_CDIP") + self.class_names = ['advertisement', 'budget', 'email', 'file_folder', 'form', 'handwritten', 'invoice', 'letter', 'memo', 'news_article', 'presentation', 'questionnaire', 'resume', 'scientific_publication', 'scientific_report', 'specification'] + + def classify(self, image): + # Convert the image tensor to a PIL Image + image = Image.fromarray((image[0].numpy() * 255).astype(np.uint8).transpose(1, 2, 0)) + + # Convert the image to the model's expected input format + inputs = self.extractor(images=image, return_tensors="pt") + + # Perform the classification + outputs = self.model(**inputs) + logits = outputs.logits + predicted_class_index = torch.argmax(logits, dim=1).item() + + # Get the class name + predicted_class_name = self.class_names[predicted_class_index] + + return (predicted_class_index, predicted_class_name) + +NODE_CLASS_MAPPINGS = { + "AgeClassifierNode": AgeClassifierNode, + "FoodCategoryClassifierNode": FoodCategoryClassifierNode, + "DocumentClassificationNode": DocumentClassificationNode, + "ImageCaptioning": ImageCaptioningNode, + "ArtOrHumanClassifierNode": ArtOrHumanClassifierNode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "ImageScaleToMin": "Image Scale To Min", + "ImageCaptioning": "Image Captioning", + "ArtOrHumanClassifierNode": "Art Or Human Classifier", +} \ No newline at end of file diff --git a/nodes/ImageProcessingNode.py b/nodes/ImageProcessingNode.py new file mode 100644 index 0000000..7a6efcf --- /dev/null +++ b/nodes/ImageProcessingNode.py @@ -0,0 +1,559 @@ +import hashlib +import fastapi +import fastapi +import torch, time +import io + +import comfy.samplers +from matplotlib import transforms +from PIL import Image, ImageFilter, ImageEnhance, ImageOps, ImageDraw, ImageChops, ImageFont +import numpy as np +import comfy.model_management as model_management +import json +import uuid +import os +from warnings import filterwarnings +import pytorch_lightning as pl +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) + + +aspect_ratios = [ + "1/1", # square + "4/3", # standard monitor + "3/2", # 35mm film + "16/9", # widescreen monitor + "21/9" # ultrawide monitor +] + +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) + mask_np = mask.cpu().numpy().astype('uint8') + mask_pil = Image.fromarray(mask_np, mode="L") + return mask_pil + + +class ImageRankingNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "score": ("INT",), + "prompt": ("STRING",), + "image_path": ("STRING",), + "json_file_path": ("STRING",) # JSON file path + }, + } + + RETURN_TYPES = () + FUNCTION = "rank_image" + CATEGORY = "LexTools/ImageProcessing/Ranking" + + def rank_image(self, score, prompt, image_path, json_file_path): + # Load JSON data from file + with open(json_file_path, 'r') as f: + data = json.load(f) + + # Check if prompt exists in data + for record in data: + if record["prompt"] == prompt: + # Prompt exists, append image path and score + record["generations"].append(image_path) + record["ranking"].append(score) + break + else: + # Prompt does not exist, create new record + new_id = str(uuid.uuid4()) # Generate a unique ID + new_record = { + "id": new_id, + "prompt": prompt, + "generations": [image_path], + "ranking": [score] + } + data.append(new_record) + + # Save updated data back to JSON file + with open(json_file_path, 'w') as f: + json.dump(data, f) + + +class ImageAspectPadNode: + + @classmethod + def INPUT_TYPES(s): + 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}), + + + }, + "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): + + 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 + + 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 + + mask = torch.ones( + (d2 + top_padding + bottom_padding, d3 + left_padding + right_padding), + dtype=torch.float32, + ) + + t = torch.zeros( + (d2, d3), + dtype=torch.float32 + ) + + if feathering > 0 and feathering * 2 < d2 and feathering * 2 < d3: + + 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}),}} + + RETURN_TYPES = ("FLOAT",) + FUNCTION = "calculate_scale" + + CATEGORY = "LexTools/ImageProcessing/upscaling" + + 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): + return { + "required": { + "score": ("INT", {"default": 0}), + "threshold": ("INT", {"default": 0}), + "image": ("IMAGE", {"default": None}), + }, + } + + 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,) + +class ImageFilterByFloatScoreNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "score": ("FLOAT", {"default": 0.0}), + "threshold": ("FLOAT", {"default": 0.0}), + "image": ("IMAGE", {"default": None}), + }, + } + + 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,) + +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}), + }, + } + OUTPUT_NODE = True + + RETURN_TYPES = ("FLOAT",) + FUNCTION = "calculate_score" + CATEGORY = "LexTools/ImageProcessing/Scores" + + 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 + + # Prepare the output UI + return (final_score, {"ui": {"STRING": [final_score]}}) + + +# +# Class taken from https://github.com/christophschuhmann/improved-aesthetic-predictor simple_inference.py +# +class MLP(pl.LightningModule): + def __init__(self, input_size, xcol='emb', ycol='avg_rating'): + super().__init__() + self.input_size = input_size + self.xcol = xcol + self.ycol = ycol + self.layers = nn.Sequential( + nn.Linear(self.input_size, 1024), + #nn.ReLU(), + nn.Dropout(0.2), + nn.Linear(1024, 128), + #nn.ReLU(), + nn.Dropout(0.2), + nn.Linear(128, 64), + #nn.ReLU(), + nn.Dropout(0.1), + nn.Linear(64, 16), + #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 + def validation_step(self, batch, batch_idx): + x = batch[self.xcol] + y = batch[self.ycol].reshape(-1, 1) + 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) + +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 CalculateAestheticScore: + device = "cuda" + model2 = None + preprocess = None + model = None + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "aesthetic_model": ("AESTHETIC_MODEL",), + }, + "optional": { + "keep_in_memory": ("BOOL", {"default": True}), + } + } + + RETURN_TYPES = ("SCORE",) + FUNCTION = "execute" + CATEGORY = "LexTools/ImageProcessing/aestheticscore" + + 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 + + m_path2 = aesthetic_model + + if not self.model: + 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] + img = (tensor_image * 255).to(torch.uint8).numpy() + pil_image = Image.fromarray(img, mode='RGB') + + # Use the class variable preprocess + image2 = self.preprocess(pil_image).unsqueeze(0).to(self.device) + + with torch.no_grad(): + # Use the class variable model2 + 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)) + final_prediction = int(float(prediction[0])*100) + + if not keep_in_memory: + self.model = None + self.model2 = None + self.preprocess = None + + return (final_prediction,) + +class MD5ImageHashNode: + device = "cuda" + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "execute" + CATEGORY = "LexTools/ImageProcessing/md5hash" + + 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') + + # Convert PIL image to bytes + image_byte_arr = io.BytesIO() + pil_image.save(image_byte_arr, format='PNG') + image_byte_arr = image_byte_arr.getvalue() + + # Calculate MD5 hash + m = hashlib.md5() + m.update(image_byte_arr) + md5_hash = m.hexdigest() + + return (md5_hash,) + +class AesthetlcScoreSorter: + def __init__(self): + pass + pass + @classmethod + def INPUT_TYPES(s): + return { + "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,) + +class ScoreConverterNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "score": ("SCORE", {"default": 0.0}), + }, + "optional": { + "show_on_node": ("INT", {"default": 0}), + }, + } + + RETURN_TYPES = ("INT", "FLOAT", "STRING",) + FUNCTION = "convert_score" + CATEGORY = "LexTools/ImageProcessing/Scores" + OUTPUT_NODE = True + + def convert_score(self, score, show_on_node): + # Convert the score to an integer, float, and string + score_int = int(score) + score_float = float(score) + score_str = str(score) + + # 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) + +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_TYPES = ("STRING","INT","FLOAT","FLOAT","STRING","STRING") + FUNCTION = "sample" + + CATEGORY = "sampling" + + def sample(self, ckpt_name, steps, cfg, sampler_name, scheduler,denoise): + pass + return (ckpt_name, steps, cfg, sampler_name, scheduler,denoise) + +NODE_CLASS_MAPPINGS = { + + "ImageFilterByIntScoreNode": ImageFilterByIntScoreNode, + "ImageFilterByFloatScoreNode": ImageFilterByFloatScoreNode, + "ImageScaleToMin": ImageScaleToMin, + "ImageAspectPadNode": ImageAspectPadNode, + "ImageRankingNode": ImageRankingNode, + "ImageQualityScoreNode": ImageQualityScoreNode, + "ScoreConverterNode":ScoreConverterNode, + "MD5ImageHashNode": MD5ImageHashNode, + "SamplerPropertiesNode": SamplerPropertiesNode, + +} + + +NODE_DISPLAY_NAME_MAPPINGS = { + "ImageFilterByIntScoreNode": "Image Filter (Int Score)", + "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":"Sampler input node", + } \ No newline at end of file diff --git a/nodes/SegformerNode.py b/nodes/SegformerNode.py new file mode 100644 index 0000000..658ec11 --- /dev/null +++ b/nodes/SegformerNode.py @@ -0,0 +1,371 @@ +import torch +from transformers import SegformerImageProcessor, AutoModelForSemanticSegmentation +from PIL import Image,ImageOps,ImageFilter +import torch.nn as nn +import matplotlib.pyplot as plt +import numpy as np +import io +from scipy.ndimage import binary_dilation + + +model_names = [ + "enes361/segformer_b2_clothes", + "mattmdjaga/segformer_b0_clothes", + "mattmdjaga/segformer_b2_clothes", + "DiTo97/binarization-segformer-b3", + "s3nh/SegFormer-b0-person-segmentation", + "venture361/clothes_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 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]}), + + + }, + } + + RETURN_TYPES = ("IMAGE","MASK", "STRING") + FUNCTION = "segment_image" + CATEGORY = "LexTools/ImageProcessing/Segmentation" + + def __init__(self): + pass + # 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) + + i = ImageOps.exif_transpose(img2) + if i.getbands() != ("R", "G", "B", "A"): + i = i.convert("RGBA") + + + img2 = np.array(img2).astype(np.float32) / 255.0 + img2 = torch.from_numpy(img2)[None,] + + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + else: + mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") + # Get the unique segments in the image + unique_segments = np.unique(pred_seg) + + # Create a string with the information for each segment + segment_info = [] + for segment in unique_segments: + # Get the name of the segment from the model's configuration + segment_name = self.model.config.id2label[segment] + + # Here, you would replace these values with the actual accuracy and IoU for the segment + + + segment_info.append(f"Segment {segment}: {segment_name}") + + # Join the segment info strings into a single string + segment_info_str = "\n".join(segment_info) + + + output_ui = {"images": [img2]} if show_on_node else {} + + return {"result": (img2,mask, segment_info_str), "ui": output_ui} + +class SegformerNodeMasks: + @classmethod + def INPUT_TYPES(cls): + global model_names # Assuming model_names is a list of model names + return { + "required": { + "image": ("IMAGE", {"default": None}), + "segments_to_merge": ("STRING", {"default": "0"}), + "model_name": (model_names, {"default": model_names[0]}) + + }, + } + + RETURN_TYPES = ("IMAGE", "MASK", "STRING") + FUNCTION = "segment_image" + CATEGORY = "LexTools/ImageProcessing/Segmentation" + + def __init__(self): + pass + + # 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): + # 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)) + inputs = self.processor(images=img, return_tensors="pt") + + # Get the outputs from the model + outputs = self.model(**inputs) + logits = outputs.logits.cpu() + + # Upsample the logits to match the original image size + upsampled_logits = nn.functional.interpolate( + logits, + size=img.size[::-1], + mode="bilinear", + align_corners=False, + ) + + # Get the predicted segments + pred_seg = upsampled_logits.argmax(dim=1)[0] + unique_segments = np.unique(pred_seg) + + # Initialize lists to hold the segmented images and masks + segmented_images = [] + masks = [] + + # Iterate over the unique segments + for segment in unique_segments: + # Create a binary mask for the current segment + mask = np.where(pred_seg == segment, 1, 0).astype(np.uint8) + # Upsample the mask to match the original image size + mask = nn.functional.interpolate(torch.from_numpy(mask)[None, None,], size=img.size[::-1], mode="nearest")[0] + + # Apply the mask to the original image to get the segmented image + segmented_image = img * mask.numpy()[0, ..., None] + segmented_image_pil = Image.fromarray(segmented_image) + + # Convert the segmented image and mask to tensors + img2 = torch.from_numpy(np.array(segmented_image_pil).astype(np.float32) / 255.0)[None,] + segmented_images.append(img2) + masks.append(mask) + + # Initialize a mask of zeros with the same size as the other masks + merged_mask = torch.zeros_like(masks[0]) + + # Iterate over the segments to merge + for segment in segments_to_merge: + # Check if the segment index is valid + if segment < len(masks): + # Add the current mask to the merged mask + merged_mask += masks[segment] + else: + # Raise an error if the segment index is invalid + raise ValueError(f"Segment {segment} is out of range. There are only {len(masks)} segments.") + + # Get the merged image by applying the merged mask to the original image + merged_image = img * merged_mask.numpy()[0, ..., None] + merged_image_pil = Image.fromarray(merged_image) + img2 = torch.from_numpy(np.array(merged_image_pil).astype(np.float32) / 255.0)[None,] + + + 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 + + }, + } + + OUTPUT_NODE = True + + RETURN_TYPES = ("IMAGE", "MASK", "STRING") + 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 + + 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) + + 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)) + + img2 = np.array(merged_image_pil).astype(np.float32) / 255.0 + img2 = torch.from_numpy(img2).double()[None,] + + merged_mask_torch = torch.from_numpy(merged_mask).float()[None,] # change from double to float + + merged_segments_str = ','.join(map(str, merged_segments)) + + output_ui = {"images": [img2]} if show_on_node else {} + + return {"result": (img2, merged_mask_torch, merged_segments_str), "ui": output_ui} + + + +class SeedIncrementerNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "seed": ("INT", {"default": 0, "min": 0,"max": 0xffffffffffffffff}), + "IncrementAt": ("INT", {"default": 10, "min": 1,"max": 0xffffffffffffffff}), + }, + } + + RETURN_TYPES = ("STRING", "INT","STRING", "INT") + FUNCTION = "increment_seed" + CATEGORY = "LexTools/Utilities" + + def increment_seed(self, seed, IncrementAt): + # Compute subseed + subseed = seed // IncrementAt + 1 + + + # Return seed as string, seed as int, and subseed + return str(seed), seed, str(subseed),subseed + + +class StepCfgIncrementNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "seed": ("INT", {"default": 0,"max": 0xffffffffffffffff}), + "cfg_start": ("INT", {"default": 7}), + "steps_start": ("INT", {"default": 10}), + "image_steps": ("INT", {"default": 100}), + "max_steps": ("INT", {"default": 12}), + }, + } + + RETURN_TYPES = ("INT", "INT") + FUNCTION = "calculate_steps_cfg" + CATEGORY = "LexTools/ImageProcessing/Increment" + + def calculate_steps_cfg(self, seed, cfg_start, steps_start, image_steps, max_steps): + # Calculate the number of complete cycles + cycle_count = seed // (image_steps * (max_steps - steps_start + 1)) + + # Calculate the number of steps within the current cycle + step_in_cycle = (seed // image_steps) % (max_steps - steps_start + 1) + + # Update cfg and steps values + cfg = cfg_start + cycle_count + steps = steps_start + step_in_cycle + + return cfg, steps + + + + +NODE_CLASS_MAPPINGS = { + "SegformerNode": SegformerNode, + "SegformerNodeMasks": SegformerNodeMasks, + "SegformerNodeMergeSegments": SegformerNodeMergeSegments, + "SeedIncrementerNode": SeedIncrementerNode, + "StepCfgIncrementNode": StepCfgIncrementNode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "SegformerNode": "Segformer Node", + "SegformerNodeMasks": "Segformer Node Masks", + "SegformerNodeMergeSegments": "Segformer Node Merge Segments", + "SeedIncrementerNode": "Seed Incrementer Node", + "StepCfgIncrementNode": "Step Cfg Increment Node", +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..82a77ad --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +numpy +opencv-python +git+https://github.com/facebookresearch/detectron2.git +pyodbc \ No newline at end of file