diff --git a/__init__.py b/__init__.py index 9566337..ec578ad 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..1db657b --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +boto3 \ No newline at end of file