Add huggingface caption models

This commit is contained in:
Yolan
2023-07-28 21:54:31 -07:00
parent 2c53b59d1b
commit 36c121c586
5 changed files with 131 additions and 69 deletions
+2 -2
View File
@@ -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.
+49 -67
View File
@@ -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"
}
+21
View File
@@ -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'))
+1
View File
@@ -0,0 +1 @@
requests
+58
View File
@@ -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'))