Add s3 uploader
This commit is contained in:
+128
-39
@@ -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
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
boto3
|
||||
Reference in New Issue
Block a user