feat: avatar generation endpoint

This commit is contained in:
Radionic
2023-11-27 12:45:38 +08:00
parent 34c5dcf65b
commit 02be90f8ff
+108 -4
View File
@@ -12,6 +12,8 @@ import re
import base64
from PIL import Image
import io
import time
import execution
load_dotenv()
@@ -102,11 +104,17 @@ async def post_sam_model(request):
"Content-type": "application/json",
"Accept": "application/json",
},
data=json.dumps({
"image": image,
}),
data=json.dumps(
{
"image": image,
}
),
).json()
emb, input_size, original_size = res["emb"], res["input_size"], res["original_size"]
emb, input_size, original_size = (
res["emb"],
res["input_size"],
res["original_size"],
)
emb = np.array(emb).astype(np.float32)
np.save(emb_filename, emb)
with open(output_json_filename, "w") as f:
@@ -135,6 +143,102 @@ async def post_sam_model(request):
return web.json_response({})
def save_image(image):
input_folder = folder_paths.get_input_directory()
name, extension = os.path.splitext(image.filename)
save_name = f"{name}{extension}"
i = 1
while os.path.exists(f"{input_folder}/{save_name}"):
save_name = f"{name}_{i}{extension}"
i += 1
with open(f"{input_folder}/{save_name}", "wb") as f:
f.write(image.file.read())
return save_name
def post_prompt(json_data):
prompt_server = server.PromptServer.instance
json_data = prompt_server.trigger_on_prompt(json_data)
if "number" in json_data:
number = float(json_data["number"])
else:
number = prompt_server.number
if "front" in json_data:
if json_data["front"]:
number = -number
prompt_server.number += 1
if "prompt" in json_data:
prompt = json_data["prompt"]
valid = execution.validate_prompt(prompt)
extra_data = {}
if "extra_data" in json_data:
extra_data = json_data["extra_data"]
if "client_id" in json_data:
extra_data["client_id"] = json_data["client_id"]
if valid[0]:
prompt_id = str(uuid.uuid4())
outputs_to_execute = valid[2]
prompt_server.prompt_queue.put(
(number, prompt_id, prompt, extra_data, outputs_to_execute)
)
response = {
"prompt_id": prompt_id,
"number": number,
"node_errors": valid[3],
}
return web.json_response(response)
else:
print("invalid prompt:", valid[1])
return web.json_response(
{"error": valid[1], "node_errors": valid[3]}, status=400
)
else:
return web.json_response({"error": "no prompt", "node_errors": []}, status=400)
def get_avatar_file(outputs):
for node_id, output in outputs.items():
if "gltfFilename" in output:
avatar_filename = output["gltfFilename"][0]
with open(
f"{folder_paths.get_output_directory()}/{avatar_filename}", "rb"
) as f:
return f.read()
with open(
os.path.join(
os.path.dirname(__file__), "workflow_templates/api/avatar_generation_api.json"
)
) as f:
api_prompt_string = "\n".join(f.readlines())
@server.PromptServer.instance.routes.post("/avatar_generation")
async def post_prompt_block(request):
prompt_server = server.PromptServer.instance
post = await request.post()
image = post.get("image")
image_name = save_image(image)
api_prompt = json.loads(api_prompt_string.replace("IMAGE_REFERENCE", image_name))
json_data = {"prompt": api_prompt}
res = post_prompt(json_data)
prompt_id = json.loads(res.text)["prompt_id"]
while True:
history = prompt_server.prompt_queue.get_history(prompt_id=prompt_id)
if history:
file = get_avatar_file(history[prompt_id]["outputs"])
return web.Response(body=file)
time.sleep(0.5)
# @server.PromptServer.instance.routes.get("/get_default_workflow")
# async def get_default_workflow(request):
# # json_link = "https://cdn.discordapp.com/attachments/1119102674437156984/1172255632586448987/workflow_boy_2_1.json?ex=655fa722&is=654d3222&hm=463fa6a3c6ea60f7471196ff45382c729d3b856e86282f905d37a0398711860e&" # YP workflow