diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..8374e77 --- /dev/null +++ b/__init__.py @@ -0,0 +1,6 @@ +from .easyapi.ImageNode import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + + +WEB_DIRECTORY = "./static" + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/easyapi/ImageNode.py b/easyapi/ImageNode.py new file mode 100644 index 0000000..4d1eb9f --- /dev/null +++ b/easyapi/ImageNode.py @@ -0,0 +1,173 @@ +import base64 +import io + +import numpy as np +import torch +from PIL import ImageOps, Image +from nodes import LoadImage +from comfy.cli_args import args +from PIL.PngImagePlugin import PngInfo +import json +from json import JSONEncoder, JSONDecoder + +class Base64ToImage: + @classmethod + def INPUT_TYPES(self): + return {"required": { + "base64Images": ("STRING", {"forceInput": True}), + }, + } + + RETURN_TYPES = ("IMAGE",) + # RETURN_NAMES = ("image", "mask") + + FUNCTION = "convert" + + CATEGORY = "easyapi/image" + + # INPUT_IS_LIST = False + OUTPUT_IS_LIST = (True, False) + + def convert(self, base64Images): + # print(base64Image) + base64ImageJson = JSONDecoder().decode(s=base64Images) + images = [] + for base64Image in base64ImageJson: + i = base64_to_image(base64Image) + # 下面代码参考LoadImage + i = ImageOps.exif_transpose(i) + image = i.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None, ] + # if 'A' in i.getbands(): + # mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + # mask = 1. - torch.from_numpy(mask) + # else: + # mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") + images.append(image) + + return torch.stack(images, dim=0)[None, ] + # return (torch.stack(images, dim=0)[None, ], mask.unsqueeze(0)) + + +class ImageToBase64: + @classmethod + def INPUT_TYPES(self): + return {"required": { + "images": ("IMAGE",), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("base64Images",) + + FUNCTION = "convert" + # 作为输出节点,返回数据格式是{"ui": {output_name:value}, "result": (value,)} + # ui中是websocket返回给前端的内容,result是py执行传给下个节点用的 + OUTPUT_NODE = True + + CATEGORY = "easyapi/image" + + # INPUT_IS_LIST = False + # OUTPUT_IS_LIST = (False,False,) + + def convert(self, images, prompt=None, extra_pnginfo=None): + result = list() + for i in images: + img = tensor_to_pil(i) + + # 创建一个BytesIO对象,用于临时存储图像数据 + image_data = io.BytesIO() + metadata = None + if not args.disable_metadata: + 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])) + + # 将图像保存到BytesIO对象中,格式为PNG + img.save(image_data, format='PNG', pnginfo=metadata) + + # 将BytesIO对象的内容转换为字节串 + image_data_bytes = image_data.getvalue() + + # 将图像数据编码为Base64字符串 + encoded_image = "data:image/png;base64," + base64.b64encode(image_data_bytes).decode('utf-8') + result.append(encoded_image) + base64Images = JSONEncoder().encode(result) + # print(images) + return {"ui": {"base64images": result}, "result": (base64Images,)} + + +class LoadImageToBase64(LoadImage): + + RETURN_TYPES = ("STRING", "IMAGE", "MASK", ) + RETURN_NAMES = ("base64Images", "IMAGE", "MASK", ) + + FUNCTION = "convert" + OUTPUT_NODE = True + + CATEGORY = "easyapi/image" + + # INPUT_IS_LIST = False + # OUTPUT_IS_LIST = (False,False,) + + def convert(self, image): + img, mask = self.load_image(image) + + i = tensor_to_pil(img) + # 创建一个BytesIO对象,用于临时存储图像数据 + image_data = io.BytesIO() + + # 将图像保存到BytesIO对象中,格式为PNG + i.save(image_data, format='PNG') + + # 将BytesIO对象的内容转换为字节串 + image_data_bytes = image_data.getvalue() + + # 将图像数据编码为Base64字符串 + encoded_image = "[\"data:image/png;base64," + base64.b64encode(image_data_bytes).decode('utf-8') + "\"]" + return encoded_image, img, mask + + +# Tensor to PIL +def tensor_to_pil(image): + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + + +# Convert PIL to Tensor +def pil_2_tensor(image): + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) + + +def base64_to_image(base64_string): + # 去除前缀 + prefix, base64_data = base64_string.split(",", 1) + + # 从base64字符串中解码图像数据 + image_data = base64.b64decode(base64_data) + + # 创建一个内存流对象 + image_stream = io.BytesIO(image_data) + + # 使用PIL的Image模块打开图像数据 + image = Image.open(image_stream) + + return image + + +NODE_CLASS_MAPPINGS = { + "Base64ToImage": Base64ToImage, + "ImageToBase64": ImageToBase64, + "LoadImageToBase64": LoadImageToBase64, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "Base64ToImage": "Base64 To Image", + "ImageToBase64": "Image To Base64", + "LoadImageToBase64": "Load Image To Base64", +} diff --git a/static/js/image_node.js b/static/js/image_node.js new file mode 100644 index 0000000..385cbf9 --- /dev/null +++ b/static/js/image_node.js @@ -0,0 +1,29 @@ +import { app } from "/scripts/app.js"; +import { api } from "/scripts/api.js"; +// ================= CREATE EXTENSION ================ +/*app.registerExtension({ + name: "Comfy.EasyApiImageNode", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData.name === "Base64ToImage") { + console.log(nodeData) + } + }, +});*/ +api.addEventListener("executed", ({detail}) => { + const images = detail?.output?.base64images; + if (!images) return; + const currentNode = app.graph._nodes_by_id[detail.node]; + // console.log(currentNode.imgs) + currentNode.imgs = []; + for(let i in images){ + let img = images[i] + let image = new Image() + image.onload = () => { + currentNode.imgs.push(image); + currentNode.setSizeForImage?.(); + app.graph.setDirtyCanvas(true, true); + }; + image.src=img; + } +}); +// ================= END CREATE EXTENSION ================