diff --git a/__init__.py b/__init__.py index afc0c33..e8ff1c4 100644 --- a/__init__.py +++ b/__init__.py @@ -15,6 +15,7 @@ sys.path.append(extension_folder) logScript.log_wrap() api.init() + def loadCustomNodes(): files = glob.glob(os.path.join(pyPath, "*Node.py"), recursive=True) for file in files: diff --git a/easyapi/ImageNode.py b/easyapi/ImageNode.py index 927d094..cc5bfc5 100644 --- a/easyapi/ImageNode.py +++ b/easyapi/ImageNode.py @@ -8,7 +8,7 @@ from comfy.cli_args import args from PIL.PngImagePlugin import PngInfo import json from json import JSONEncoder, JSONDecoder -from easyapi.util import tensor_to_pil +from easyapi.util import tensor_to_pil, base64_to_image, image_to_base64 class Base64ToImage: @@ -87,9 +87,6 @@ class ImageToBase64Advanced: 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() @@ -103,14 +100,8 @@ class ImageToBase64Advanced: 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') + encoded_image = image_to_base64(img, pnginfo=metadata) result.append(encoded_image) base64Images = JSONEncoder().encode(result) # print(images) @@ -254,22 +245,6 @@ class LoadImageToBase64(LoadImage): return encoded_image, img, mask -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, diff --git a/easyapi/api.py b/easyapi/api.py index f5f1ef0..ad4d7e1 100644 --- a/easyapi/api.py +++ b/easyapi/api.py @@ -1,13 +1,24 @@ import os import json +import folder_paths import nodes from server import PromptServer from aiohttp import web import execution +from simple_lama_inpainting import SimpleLama +from .util import image_to_base64, base64_to_image extension_folder = os.path.dirname(os.path.realpath(__file__)) +simple_lama = None +lama_model_path = os.path.join(folder_paths.models_dir, "lama/big-lama.pt") +if not os.path.exists(lama_model_path): + os.environ['LAMA_MODEL'] = '' + print(f"## lama model not found: {lama_model_path}, pls download from https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt") +else: + os.environ['LAMA_MODEL'] = lama_model_path + def reset_history_size(max_size=execution.MAXIMUM_HISTORY_SIZE, isStart=False): configDataFilePath = os.path.join(extension_folder, 'config') @@ -116,6 +127,29 @@ def register_routes(): PromptServer.instance.prompt_queue.delete_queue_item(delete_func) return web.Response(status=200) + @PromptServer.instance.routes.post("/easyapi/lama_cleaner") + async def lama_cleaner(request): + json_data = await request.json() + image = json_data["image"] + mask = json_data["mask"] + if image is None or mask is None: + return web.json_response({"error": "missing required params"}, status=400) + + global simple_lama + if simple_lama is None: + simple_lama = SimpleLama() + + image = base64_to_image(image) + mask = base64_to_image(mask) + mask = mask.convert('L') + + res = simple_lama(image, mask) + + encoded_image = image_to_base64(res) + + response = {"base64Image": encoded_image} + return web.json_response(response, status=200) + def init(): reset_history_size(isStart=True) diff --git a/easyapi/util.py b/easyapi/util.py index af91268..be9d594 100644 --- a/easyapi/util.py +++ b/easyapi/util.py @@ -1,7 +1,11 @@ +import base64 +import io + import numpy as np import torch from PIL import Image + # Tensor to PIL def tensor_to_pil(image): return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) @@ -11,3 +15,38 @@ def tensor_to_pil(image): 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): + # 去除前缀 + base64_list = base64_string.split(",", 1) + if len(base64_list) == 2: + prefix, base64_data = base64_list + else: + base64_data = base64_list[0] + + # 从base64字符串中解码图像数据 + image_data = base64.b64decode(base64_data) + + # 创建一个内存流对象 + image_stream = io.BytesIO(image_data) + + # 使用PIL的Image模块打开图像数据 + image = Image.open(image_stream) + + return image + + +def image_to_base64(pli_image, pnginfo=None): + # 创建一个BytesIO对象,用于临时存储图像数据 + image_data = io.BytesIO() + + # 将图像保存到BytesIO对象中,格式为PNG + pli_image.save(image_data, format='PNG', pnginfo=pnginfo) + + # 将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 diff --git a/requirements.txt b/requirements.txt index 4032d22..156036f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -segment_anything \ No newline at end of file +segment_anything +simple_lama_inpainting \ No newline at end of file