From 5273e07eacd790c4cec2e87417794d26289dc20e Mon Sep 17 00:00:00 2001 From: Radionic Date: Wed, 6 Dec 2023 15:52:16 +0800 Subject: [PATCH] feat: support base image as input --- routes.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/routes.py b/routes.py index 0ddf618..d4b1706 100644 --- a/routes.py +++ b/routes.py @@ -253,7 +253,7 @@ def randomSeed(num_digits=15): with open( os.path.join( os.path.dirname(__file__), - "workflow_templates/api/avatar_generation_mask_api.json", + "workflow_templates/api/avatar_generation_mask_api_v2.json", ) ) as f: default_workflow = "\n".join(f.readlines()) @@ -266,18 +266,26 @@ async def post_prompt_block(request): workflow = post.get("workflow") workflow = default_workflow if workflow is None else workflow - image = post.get("image") - image_path = save_image(image) - image_name, image_ext = os.path.splitext(image_path) - workflow = workflow.replace("IMAGE_REFERENCE", image_path).replace( - "SEED", str(randomSeed()) - ) + ref_image = post.get("ref_image") + base_image = post.get("base_image") + if ref_image is not None: + image_path = save_image(ref_image) + image_name, image_ext = os.path.splitext(image_path) + workflow = workflow.replace("reference_image", image_path) + elif base_image is not None: + image_path = save_image(base_image) + image_name, image_ext = os.path.splitext(image_path) + workflow = workflow.replace("base_image", image_path) + workflow = workflow.replace("reference_image", image_path) # TMP + for key, value in post.items(): if key.startswith("mask_"): mask_name = image_name + "_" + key.replace("mask_", "") + image_ext mask_path = save_image(value, save_name=mask_name) workflow = workflow.replace(key, mask_path) + workflow = workflow.replace("SEED", str(randomSeed())) + api_prompt = json.loads(workflow) res = post_prompt({"prompt": api_prompt}) prompt_id = json.loads(res.text)["prompt_id"]