add lama api

This commit is contained in:
刘雪峰
2024-01-16 16:31:29 +08:00
parent 1a8c571646
commit 499045cc84
5 changed files with 78 additions and 28 deletions
+2 -27
View File
@@ -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,
+34
View File
@@ -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)
+39
View File
@@ -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