feat: support base image as input

This commit is contained in:
Radionic
2023-12-06 15:52:30 +08:00
parent 51ce8ba957
commit 5273e07eac
+15 -7
View File
@@ -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"]