Add s3 uploader

This commit is contained in:
Yolan
2023-07-12 01:17:33 -07:00
parent e9cbcfd5d8
commit f2649cdd0b
2 changed files with 129 additions and 39 deletions
+128 -39
View File
@@ -1,5 +1,9 @@
import uuid
import torch
import boto3
import os
import sys
import json
@@ -32,6 +36,8 @@ from urllib import request
from PIL import Image, ImageOps
import folder_paths
from comfy_extras.chainner_models import model_loading
from custom_nodes.DTGlobalVariables import variables
try:
from torchvision.transforms import ToPILImage
@@ -107,7 +113,7 @@ class RemoteLoader:
return filename
upscalers = RemoteLoader("checkpoints", "https://api.aiart.doubtech.com/comfyui/upscalers")
checkpoints = RemoteLoader("checkpoints", "https://api.aiart.doubtech.com/comfyui/checkpoints")
vae = RemoteLoader("vae", "https://api.aiart.doubtech.com/comfyui/vae")
lora = RemoteLoader("loras", "https://api.aiart.doubtech.com/comfyui/lora")
@@ -170,10 +176,10 @@ class SubmitImage:
"""
return {
"required": {
"image": ("IMAGE",),
"images": ("IMAGE",),
},
"optional": {
"prompt": ("STRING", {
"prompt_text": ("STRING", {
"multiline": False, # True if you want the field to look like the one on the ClipTextEncode node
"default": ""
}),
@@ -201,6 +207,7 @@ class SubmitImage:
"default": False
})
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ()
@@ -210,41 +217,108 @@ class SubmitImage:
OUTPUT_NODE = True
CATEGORY = "DoubTech/image"
CATEGORY = "DoubTech/Image"
def upload(self, image, prompt, tags, title, alt, caption, set, private):
def upload_image_to_s3(self, image, pnginfo=None):
# Create an S3 client
s3_client = boto3.client('s3', aws_access_key_id=config.s3_key, aws_secret_access_key=config.s3_secret, region_name=config.s3_region)
# Convert the image to bytes
img_bytes = io.BytesIO()
image.save(img_bytes, format='PNG', pnginfo=pnginfo)
img_bytes.seek(0)
# Generate a UUID-based name for the S3 key
s3_key = f"images/{uuid.uuid4()}.png"
# Upload the image to the S3 bucket
try:
s3_client.upload_fileobj(img_bytes, config.s3_bucket, s3_key)
print(f"Image uploaded to S3 bucket: {config.s3_bucket}, key: {s3_key}")
# Generate the URL of the uploaded image
s3_url = f"https://{config.s3_bucket}.s3.amazonaws.com/{s3_key}"
return s3_url
except Exception as e:
print("Error uploading image to S3:", e)
return None
def upload(self, images, prompt_text, tags, title, alt, caption, set, private, prompt=None, extra_pnginfo=None):
print("uploading image...")
if variables.generated_prompt is not None:
prompt_text = variables.generated_prompt
# uriencode the parameters
tags = urllib.parse.quote(tags)
title = urllib.parse.quote(title)
alt = urllib.parse.quote(alt)
setName = urllib.parse.quote(set)
prompt = urllib.parse.quote(prompt)
caption = urllib.parse.quote(caption)
if "tags" in variables.state:
current_tags = variables.state["tags"].split(",")
current_tags = [t for t in current_tags if t != ""]
tags = tags.split(",") or []
tags = [t for t in tags if t != ""]
# Merge the two arrays into one and remove duplicates
tags = current_tags + [x for x in tags if x not in current_tags]
tags = ",".join(tags)
tags = urllib.parse.quote(variables.apply(tags, "tags"))
title = urllib.parse.quote(variables.apply(title, "title"))
alt = urllib.parse.quote(variables.apply(alt, "alt"))
setName = urllib.parse.quote(variables.apply(set, "set"))
prompt_text = urllib.parse.quote(variables.apply(prompt_text))
caption = urllib.parse.quote(variables.apply(caption))
# Create a post request to submit the image as the post body to the backend
uri = "https://api.aiart.doubtech.com/comfyui/submit?key={}&tags={}&title={}&alt={}&set={}&prompt={}&caption={}&private={}".format(config.apikey, tags, title, alt, setName, prompt, caption, private)
uri = "https://api.aiart.doubtech.com/comfyui/submit?key={}&tags={}&title={}&alt={}&set={}&prompt={}&caption={}&private={}".format(
config.apikey,
tags,
title,
alt,
setName,
prompt_text,
caption,
private)
print(f"Submitting {prompt_text} with data:\n{prompt}")
num_images = image.size(0)
#iterate over the images
for i in range(num_images):
print("There are " + str(num_images) + " images in the batch")
img = image[i]
for image in images:
# Convert the image to a png
png = ToPILImage()(img.permute(2, 0, 1))
img_bytes = io.BytesIO()
png.save(img_bytes, format='PNG')
img_bytes.seek(0)
i = 255. * image.cpu().numpy()
png = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
metadata = PngInfo()
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
# Encode the image bytes to base64
encoded = base64.b64encode(img_bytes.getvalue()).decode('utf-8')
width = png.width
height = png.height
uri += "&width={}&height={}".format(width, height)
print("Submission uri: " + uri)
# Submit the image using a POST request
with request.urlopen(request.Request(uri, data=encoded.encode('utf-8'), method='POST')) as f:
print(f.read().decode('utf-8'))
if config.use_s3:
# Upload the image to S3
s3_url = self.upload_image_to_s3(png, pnginfo=metadata)
if s3_url is None:
return ()
uri += "&image_url=" + urllib.parse.quote(s3_url)
print("Submission uri: " + uri)
print(" (S3 URL: " + s3_url + ")")
with request.urlopen(request.Request(uri, method='POST'),
timeout=600) as f:
print(f.read().decode('utf-8'))
else:
img_bytes = io.BytesIO()
png.save(img_bytes, format='PNG', pnginfo=metadata)
img_bytes.seek(0)
# Encode the image bytes to base64
encoded = base64.b64encode(img_bytes.getvalue()).decode('utf-8')
print("Submission uri: " + uri)
# Submit the image using a POST request
with request.urlopen(request.Request(uri, data=encoded.encode('utf-8'), method='POST'), timeout=600) as f:
print(f.read().decode('utf-8'))
return ()
@@ -252,12 +326,11 @@ class SubmitImage:
class DTNodeCheckpointLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "ckpt_name": (checkpoints.list(), ),
}}
return {"required": { "ckpt_name": (checkpoints.list(), ), }}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
FUNCTION = "load_checkpoint"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_checkpoint(self, ckpt_name, output_vae=True, output_clip=True):
ckpt_path = checkpoints.download(ckpt_name)
@@ -273,7 +346,7 @@ class DTVAELoader:
RETURN_TYPES = ("VAE",)
FUNCTION = "load_vae"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
#TODO: scale factor?
def load_vae(self, vae_name):
@@ -293,7 +366,7 @@ class DTLoraLoader:
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "load_lora"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_lora(self, model, clip, lora_name, strength_model, strength_clip):
if strength_model == 0 and strength_clip == 0:
@@ -312,7 +385,7 @@ class DTCLIPLoader:
RETURN_TYPES = ("CLIP",)
FUNCTION = "load_clip"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_clip(self, clip_name):
clip_path = clip.download(clip_name)
@@ -328,7 +401,7 @@ class DTCLIPVisionLoader:
RETURN_TYPES = ("CLIP_VISION",)
FUNCTION = "load_clip"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_clip(self, clip_name):
clip_path = clipVision.download(clip_name)
@@ -344,7 +417,7 @@ class DTStyleModelLoader:
RETURN_TYPES = ("STYLE_MODEL",)
FUNCTION = "load_style_model"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_style_model(self, style_model_name):
style_model_path = style.download(style_model_name)
@@ -360,7 +433,7 @@ class DTGLIGENLoader:
RETURN_TYPES = ("GLIGEN",)
FUNCTION = "load_gligen"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_gligen(self, gligen_name):
gligen_path = gligen.download(gligen_name)
@@ -376,7 +449,7 @@ class DTControlNetLoader:
RETURN_TYPES = ("CONTROL_NET",)
FUNCTION = "load_controlnet"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_controlnet(self, control_net_name):
controlnet_path = controlNet.download(control_net_name)
@@ -393,7 +466,7 @@ class DTDiffControlNetLoader:
RETURN_TYPES = ("CONTROL_NET",)
FUNCTION = "load_controlnet"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_controlnet(self, model, control_net_name):
controlnet_path = controlNetDiff.download(control_net_name)
@@ -409,7 +482,7 @@ class DTunCLIPCheckpointLoader:
RETURN_TYPES = ("MODEL", "CLIP", "VAE", "CLIP_VISION")
FUNCTION = "load_checkpoint"
CATEGORY = "DoubTech/loaders"
CATEGORY = "DoubTech/Loaders"
def load_checkpoint(self, ckpt_name, output_vae=True, output_clip=True):
ckpt_path = unclipCheckpoint.download(ckpt_name)
@@ -589,6 +662,21 @@ class DTLoadImageMask:
return True
class DTUpscaleModelLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model_name": (upscalers.list(), ), }}
RETURN_TYPES = ("UPSCALE_MODEL",)
FUNCTION = "load_model"
CATEGORY = "DoubTech/Loaders"
def load_model(self, model_name):
model_path = upscalers.download(model_name)
sd = comfy.utils.load_torch_file(model_path, safe_load=True)
out = model_loading.load_state_dict(sd).eval()
return (out, )
# A dictionary that contains all nodes you want to export with their names
# NOTE: names should be globally unique
@@ -609,6 +697,7 @@ NODE_CLASS_MAPPINGS = {
"DTLoadLatent": DTLoadLatent,
"DTLoadImage": DTLoadImage,
"DTLoadImageMask": DTLoadImageMask,
"DTUpscaleModelLoader": DTUpscaleModelLoader,
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
+1
View File
@@ -0,0 +1 @@
boto3