From 36c121c586c0d824254c4799c5c9d3f91fc23029 Mon Sep 17 00:00:00 2001 From: Yolan Date: Fri, 28 Jul 2023 21:54:31 -0700 Subject: [PATCH] Add huggingface caption models --- README.md | 4 +- __init__.py | 116 ++++++++++++++++++++-------------------------- imagetotext.py | 21 +++++++++ requirements.txt | 1 + vitimagetotext.py | 58 +++++++++++++++++++++++ 5 files changed, 131 insertions(+), 69 deletions(-) create mode 100644 imagetotext.py create mode 100644 requirements.txt create mode 100644 vitimagetotext.py diff --git a/README.md b/README.md index 011c239..77664e2 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,5 @@ -# Comfy UI Node Template -This is a template for creating custom nodes for the Comfy UI stable diffusion client. +# Image to Text Node +A ComfyAI node to convert an image to text ## Description This Python script is an optional add-on to the Comfy UI stable diffusion client. It introduces quality of life improvements by providing variable nodes and shared global variables. diff --git a/__init__.py b/__init__.py index 175202a..a76c59f 100644 --- a/__init__.py +++ b/__init__.py @@ -1,95 +1,77 @@ -class Example: - """ - A example node +from custom_nodes.DTAIImageToTextNode.imagetotext import image_url_to_text, image_to_text - Class methods - ------------- - INPUT_TYPES (dict): - Tell the main program input parameters of nodes. - Attributes - ---------- - RETURN_TYPES (`tuple`): - The type of each element in the output tulple. - RETURN_NAMES (`tuple`): - Optional: The name of each output in the output tulple. - FUNCTION (`str`): - The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() - OUTPUT_NODE ([`bool`]): - If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. - The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. - Assumed to be False if not present. - CATEGORY (`str`): - The category the node should appear in the UI. - execute(s) -> tuple || None: - The entry point method. The name of this method must be the same as the value of property `FUNCTION`. - For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. - """ +class DTAIImageUrlToTextNode: + def __init__(self): + self.url = None + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "url": ("STRING", { + "multiline": False, # True if you want the field to look like the one on the ClipTextEncode node + "default": "https://doubtech.ai/img/logo.png" + }), + }, + } + + RETURN_TYPES = ("STRING",) + # RETURN_NAMES = ("image_output_name",) + + FUNCTION = "imagetotext" + + # OUTPUT_NODE = False + + CATEGORY = "DoubTech/Image/Image To Text" + + @classmethod + def IS_CHANGED(self, url): + return self.url != url + + def imagetotext(self, url): + self.url = url + caption = image_url_to_text(url) + print("Image appears to be: " + caption) + return (caption,) + + +class DTAIImageToTextNode: def __init__(self): pass @classmethod def INPUT_TYPES(s): - """ - Return a dictionary which contains config for all input fields. - Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". - Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. - The type can be a list for selection. - - Returns: `dict`: - - Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required` - - Value input_fields (`dict`): Contains input fields config: - * Key field_name (`string`): Name of a entry-point method's argument - * Value field_config (`tuple`): - + First value is a string indicate the type of field or a list for selection. - + Secound value is a config for type "INT", "STRING" or "FLOAT". - """ return { "required": { - "image": ("IMAGE",), - "int_field": ("INT", { - "default": 0, - "min": 0, #Minimum value - "max": 4096, #Maximum value - "step": 64 #Slider's step - }), - "float_field": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), - "print_to_screen": (["enable", "disable"],), - "string_field": ("STRING", { - "multiline": False, #True if you want the field to look like the one on the ClipTextEncode node - "default": "Hello World!" - }), + "image": ("IMAGE",) }, } - RETURN_TYPES = ("IMAGE",) + RETURN_TYPES = ("STRING",) #RETURN_NAMES = ("image_output_name",) - FUNCTION = "test" + FUNCTION = "imagetotext" #OUTPUT_NODE = False - CATEGORY = "Example" + CATEGORY = "DoubTech/Image/Image To Text" - def test(self, image, string_field, int_field, float_field, print_to_screen): - if print_to_screen == "enable": - print(f"""Your input contains: - string_field aka input text: {string_field} - int_field: {int_field} - float_field: {float_field} - """) - #do some processing on the image, in this example I just invert it - image = 1.0 - image - return (image,) + def imagetotext(self, image): + caption = image_to_text(image) + print("Image appears to be: " + caption) + return (caption,) # A dictionary that contains all nodes you want to export with their names # NOTE: names should be globally unique NODE_CLASS_MAPPINGS = { - "Example": Example + "DTAIImageToTextNode": DTAIImageToTextNode, + "DTAIImageUrlToTextNode": DTAIImageUrlToTextNode, } # A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { - "Example": "Example Node" + "DTAIImageToTextNode": "Image to Text", + "DTAIImageUrlToTextNode": "Image URL to Text" } diff --git a/imagetotext.py b/imagetotext.py new file mode 100644 index 0000000..4b49941 --- /dev/null +++ b/imagetotext.py @@ -0,0 +1,21 @@ +import requests +from PIL import Image +from transformers import BlipProcessor, BlipForConditionalGeneration + +processor = BlipProcessor.from_pretrained("Salesforce/blip-image-captioning-large") +model = BlipForConditionalGeneration.from_pretrained("Salesforce/blip-image-captioning-large").to("cuda") + +def image_url_to_text(img_url): + raw_image = Image.open(requests.get(img_url, stream=True).raw).convert('RGB') + return image_to_text(raw_image) + +def image_to_text(raw_image): + # unconditional image captioning + inputs = processor(raw_image, return_tensors="pt").to("cuda") + + out = model.generate(**inputs) + return processor.decode(out[0], skip_special_tokens=True) + +# if main +if __name__ == "__main__": + print(image_url_to_text('https://doubtech-aiart.s3.amazonaws.com/images/db50fb5c-650e-4723-b5e9-22c7d81e662c.png')) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..663bd1f --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +requests \ No newline at end of file diff --git a/vitimagetotext.py b/vitimagetotext.py new file mode 100644 index 0000000..356c567 --- /dev/null +++ b/vitimagetotext.py @@ -0,0 +1,58 @@ +from io import BytesIO + +import torch +from PIL import Image, ImageOps +from transformers import pipeline +from transformers import ViTFeatureExtractor, AutoTokenizer, VisionEncoderDecoderModel +import requests + +vit_gpt2_img_caption = pipeline("image-to-text", model="nlpconnect/vit-gpt2-image-captioning") + + +def vit_gpt2_img_caption_from_url(url): + return vit_gpt2_img_caption(url)[0]['generated_text'] + + +def load_image_from_url(url): + try: + # Send a GET request to fetch the image data + response = requests.get(url) + + # Check if the request was successful + response.raise_for_status() + + # Read the image data and create a PIL image + image = Image.open(BytesIO(response.content)) + + return image + + except requests.exceptions.RequestException as e: + print(f"Error loading image from URL: {url}") + print(e) + return None + +# Model +model_id = "nttdataspain/vit-gpt2-stablediffusion2-lora" +model = VisionEncoderDecoderModel.from_pretrained(model_id) +tokenizer = AutoTokenizer.from_pretrained(model_id) +feature_extractor = ViTFeatureExtractor.from_pretrained(model_id) + +# Predict function +def predict_prompts(list_images, max_length=16): + model.eval() + pixel_values = feature_extractor(images=list_images, return_tensors="pt").pixel_values + with torch.no_grad(): + output_ids = model.generate(pixel_values, max_length=max_length, num_beams=4, return_dict_in_generate=True).sequences + + preds = tokenizer.batch_decode(output_ids, skip_special_tokens=True) + preds = [pred.strip() for pred in preds] + return preds + +def predict_prompts_from_url(url, max_length=256): + img = load_image_from_url(url) + return predict_prompts([img], max_length=256) + + +if __name__ == "__main__": + print(predict_prompts_from_url('https://doubtech-aiart.s3.amazonaws.com/images/db50fb5c-650e-4723-b5e9-22c7d81e662c.png')) + #print(image_url_to_text('https://doubtech-aiart.s3.amazonaws.com/images/db50fb5c-650e-4723-b5e9-22c7d81e662c.png')) \ No newline at end of file