Add huggingface caption models
This commit is contained in:
@@ -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
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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'))
|
||||
@@ -0,0 +1 @@
|
||||
requests
|
||||
@@ -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'))
|
||||
Reference in New Issue
Block a user