add lama api
This commit is contained in:
+2
-27
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user