mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
314 lines
12 KiB
Python
314 lines
12 KiB
Python
import requests
|
|
import argparse
|
|
import os
|
|
import json
|
|
import time
|
|
from typing import Dict, Any, Optional
|
|
from jsonpath_ng import jsonpath, parse
|
|
import shutil
|
|
from pathlib import Path
|
|
from enum import IntEnum, auto
|
|
import os
|
|
import subprocess
|
|
import signal
|
|
|
|
SERVER_IP = os.environ.get("SERVER_IP", "127.0.0.1:8188")
|
|
SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 120))
|
|
|
|
TIMEOUT_SECONDS = int(os.environ.get("TIMEOUT_SECONDS", 1000))
|
|
PROMPT_ENDPOINT = f"http://{SERVER_IP}/prompt"
|
|
V2V_WORKFLOWS_DIRECTORY = os.environ.get("V2V_WORKFLOWS_DIRECTORY",
|
|
"/workflow_configs")
|
|
|
|
|
|
COMFY_ROOT_DIRECTORY = os.environ.get("COMFY_ROOT_DIRECTORY",
|
|
"/app/ComfyUI")
|
|
|
|
|
|
BASE_LIVE_PORTRAIT_WORKFLOW = os.environ.get("BASE_LIVE_PORTRAIT_WORKFLOW", "25-07-2024/yae_LivingPortrait(Cuda)_25-07_API.json")
|
|
|
|
class PipelineType(IntEnum):
|
|
BASE = auto()
|
|
|
|
|
|
def ensure_input_directory_empty():
|
|
input_directory = Path(COMFY_ROOT_DIRECTORY) / "input"
|
|
for file in input_directory.iterdir():
|
|
if file.is_file():
|
|
file.unlink()
|
|
|
|
def ensure_output_directory_empty():
|
|
output_directory = Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid"
|
|
for file in output_directory.iterdir():
|
|
if file.is_file():
|
|
print('deleteing file: ' + str(file) )
|
|
file.unlink()
|
|
|
|
def ensure_pipeline_input_present(source_path, input_name=None):
|
|
left = Path(source_path)
|
|
left_filename = left.name
|
|
right_filename = left_filename
|
|
if input_name is not None:
|
|
right_filename = input_name
|
|
right = Path(COMFY_ROOT_DIRECTORY) / "input" / right_filename
|
|
shutil.copy(left, right)
|
|
|
|
def primary_output_directory():
|
|
return Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid"
|
|
|
|
def expected_primary_output_name(primary_output_name="LivePortrait_00001.mp4"):
|
|
directory = primary_output_directory()
|
|
output = directory / primary_output_name
|
|
return output
|
|
|
|
def ensure_output_present_at(new_filename, primary_output_name="LivePortrait_00001.mp4"):
|
|
# Comfy will create a copy of the video that includes audio, but only if the driver video had audio.
|
|
output = expected_primary_output_name(primary_output_name)
|
|
output_with_audio = primary_output_directory() / "LivePortrait_00001-audio.mp4"
|
|
if os.path.exists(output_with_audio):
|
|
output = output_with_audio
|
|
shutil.copy(output, new_filename)
|
|
|
|
def parse_args():
|
|
parser = argparse.ArgumentParser(description='Run Comfy')
|
|
parser.add_argument('--portrait-media-filename', type=str, help='path of mp4 depth from engine for preprocessing', required=True)
|
|
parser.add_argument('--driver-media-filename', type=str, help='path of mp4 outline from engine for preprocessing', required=True)
|
|
parser.add_argument('--input-is-image', action='store_true', help='input is image', required=True)
|
|
parser.add_argument('--skip-comfy-startup', action='store_true', help='controls if this script tries to start comfy', default=True)
|
|
parser.add_argument('--tmpdir', type=str, help='path of tmpdir for preprocessing', required=False)
|
|
parser.add_argument('--output-filename', type=str, help='path of output filename', required=True)
|
|
|
|
|
|
parser.add_argument('--server-restart-trigger-file', type=str, help='server restart trigger file', required=False)
|
|
return parser.parse_args()
|
|
|
|
args = parse_args()
|
|
|
|
def get_queued_jobs_count(server_address):
|
|
try:
|
|
response = requests.get(f"http://{server_address}/queue", timeout=5)
|
|
except requests.exceptions.BaseException:
|
|
print("Error while checking queue")
|
|
return None
|
|
queue_json = response.json()
|
|
queued_jobs_count = len(queue_json.get("queue_pending", []))
|
|
running_jobs_count = len(queue_json.get("queue_running", []))
|
|
return queued_jobs_count + running_jobs_count
|
|
|
|
def queued_count_is_zero(server_address):
|
|
count = get_queued_jobs_count(server_address)
|
|
print(f"Number of queued jobs detected: {count}")
|
|
if count is None:
|
|
return False
|
|
return count == 0
|
|
|
|
def trigger_comfy_restart(restart_trigger_file):
|
|
print(f"Touching {restart_trigger_file}")
|
|
with open(restart_trigger_file, "w") as f:
|
|
f.write(time.ctime())
|
|
print("Restart triggered")
|
|
|
|
# KS: we're making the retries configurable but we should not increase the retries for all workflows without confirming we'd still meet slas. Some jobs are better to just fail. Retries compound over inference-job attempt_counts, comfy startup times, and actual inference times.
|
|
def ensure_server_is_healthy(server_address, restart_trigger_file=None, retries=1):
|
|
wait_for_comfy_startup(server_address)
|
|
if queued_count_is_zero(server_address):
|
|
print("Server is healthy")
|
|
return
|
|
if restart_trigger_file is None:
|
|
print("No restart trigger file provided. Exiting.")
|
|
exit(1)
|
|
|
|
for _ in range(retries):
|
|
trigger_comfy_restart(restart_trigger_file)
|
|
wait_for_comfy_startup(server_address)
|
|
number_of_queued_jobs = get_queued_jobs_count(server_address)
|
|
if queued_count_is_zero(server_address):
|
|
print("Server is healthy")
|
|
return
|
|
print(f"Server is not healthy. Number of queued jobs: {number_of_queued_jobs}")
|
|
print("Server is not healthy after restart. Retrying restart")
|
|
|
|
|
|
def interrupt_comfy_queue(server_address):
|
|
requests.post(f"http://{server_address}/interrupt")
|
|
print("Queue interrupted successfully.")
|
|
|
|
def wait_for_comfy_startup(server_address, timeout=SERVER_START_TIMEOUT):
|
|
prompt_endpoint = f"http://{server_address}/prompt"
|
|
start_time = time.time()
|
|
started = False
|
|
request_timeout = 2
|
|
while time.time() - start_time < timeout:
|
|
try:
|
|
response = requests.get(prompt_endpoint, timeout=request_timeout)
|
|
if response.status_code == 200:
|
|
print("Server is healthy")
|
|
started = True
|
|
break
|
|
except requests.exceptions.ConnectionError:
|
|
print(f"Server is not up yet. Time elapsed: {time.time() - start_time:.2f}")
|
|
time.sleep(1)
|
|
if not started:
|
|
print(f"Server failed to start in {SERVER_START_TIMEOUT} seconds")
|
|
exit(1)
|
|
|
|
|
|
def main():
|
|
outer_retries = 2
|
|
restart_trigger_file = args.server_restart_trigger_file
|
|
|
|
pipeline_type = PipelineType.BASE
|
|
workflow_filename = BASE_LIVE_PORTRAIT_WORKFLOW
|
|
input_media_filename = args.portrait_media_filename
|
|
input_driver_filename = args.driver_media_filename
|
|
input_is_image = args.input_is_image
|
|
|
|
for _ in range(outer_retries):
|
|
ensure_server_is_healthy(SERVER_IP, restart_trigger_file)
|
|
|
|
ensure_input_directory_empty()
|
|
ensure_output_directory_empty()
|
|
input_filename = "input.mp4"
|
|
|
|
if input_media_filename:
|
|
if input_is_image:
|
|
ensure_pipeline_input_present(input_media_filename)
|
|
input_filename = Path(input_media_filename).name
|
|
else:
|
|
ensure_pipeline_input_present(input_media_filename, "input.mp4")
|
|
|
|
if input_driver_filename:
|
|
ensure_pipeline_input_present(input_driver_filename, "driver.mp4")
|
|
print("Generating prompt json using workflow: " + workflow_filename)
|
|
prompt = generate_prompt_for_style(
|
|
workflow_filename=workflow_filename,
|
|
input_media_filename="input/" + input_filename,
|
|
input_driver_filename="input/driver.mp4",
|
|
input_is_image=input_is_image,
|
|
)
|
|
time_before = time.perf_counter()
|
|
print("running pipeline" + str(pipeline_type))
|
|
enqueue_prompt_and_wait(prompt)
|
|
time_after = time.perf_counter()
|
|
primary_output_name = "LivePortrait_00001.mp4"
|
|
if expected_primary_output_name(primary_output_name).exists():
|
|
ensure_output_present_at(args.output_filename)
|
|
print(f"Comfy pipeline ran successfully in {time_after - time_before:.2f} seconds")
|
|
# We've successfully run the pipeline, so we can exit the loop
|
|
return
|
|
else:
|
|
if restart_trigger_file is None:
|
|
print("No restart trigger file provided. Exiting.")
|
|
exit(1)
|
|
trigger_comfy_restart(restart_trigger_file)
|
|
# If we're here, we've failed to run the pipeline twice. We've tried a server restart, and we're out of retries. Maybe this was a bad input, or maybe the server is down. Either way, it's okay to give up now. We meant well.
|
|
exit(1)
|
|
|
|
def enqueue_prompt_and_wait(prompt):
|
|
data = json.dumps(prompt).encode('utf-8')
|
|
# Send HTTP post request
|
|
response = requests.post(PROMPT_ENDPOINT, data=data)
|
|
if response.status_code == 200:
|
|
print("Job enqueued successfully")
|
|
else:
|
|
print("Job failed to queue")
|
|
print(f"Status code: {response.status_code}")
|
|
print(f"Response: {response.text}")
|
|
exit(1)
|
|
|
|
# Wait for the job to finish
|
|
print(f"Waiting for job to finish (timeout: {TIMEOUT_SECONDS} seconds)")
|
|
start_time = time.time()
|
|
# check queue endpoint in loop
|
|
while True:
|
|
try:
|
|
response = requests.get(PROMPT_ENDPOINT, timeout=5)
|
|
queue_json = response.json()
|
|
if queue_json["exec_info"]["queue_remaining"]:
|
|
pass
|
|
else:
|
|
print("Job is done")
|
|
break
|
|
except Exception:
|
|
print("Error while checking queue")
|
|
time.sleep(1)
|
|
continue
|
|
if time.time() - start_time > TIMEOUT_SECONDS:
|
|
print("Timeout reached and job didn't finish!")
|
|
exit(1)
|
|
time.sleep(1)
|
|
|
|
|
|
def generate_prompt_for_style(
|
|
workflow_filename: str,
|
|
input_media_filename: Optional[str] = None,
|
|
input_driver_filename: Optional[str] = None,
|
|
input_is_image: Optional[bool] = False,
|
|
):
|
|
workflow_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "workflows"
|
|
|
|
|
|
if workflow_filename:
|
|
workflow_json_filename = workflow_filename
|
|
|
|
workflow_json = json.load(open(workflow_directory / workflow_json_filename))
|
|
|
|
|
|
jsonpath_mods = {
|
|
"$.25.inputs.video": input_driver_filename
|
|
}
|
|
|
|
if input_is_image:
|
|
jsonpath_mods["$.43.inputs.image_path"] = input_media_filename
|
|
jsonpath_mods["$.26.inputs.video"] = input_driver_filename
|
|
jsonpath_mods["$.46.inputs.boolean"] = 0
|
|
else:
|
|
jsonpath_mods["$.26.inputs.video"] = input_media_filename
|
|
jsonpath_mods["$.46.inputs.boolean"] = 1
|
|
|
|
print(*jsonpath_mods.items(), sep="\n")
|
|
|
|
json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json)
|
|
|
|
p = {"prompt": json_prompt}
|
|
return p
|
|
|
|
def apply_jsonpath_mods(jsonpath_mods: Dict[str, Any], workflow_json: Dict[str, Any]) -> Dict[str, Any]:
|
|
for key, value in jsonpath_mods.items():
|
|
# jsonpath-ng expects numbers to be wrapped in quotes
|
|
parts = key.split(".")
|
|
new_parts = [f"'{part}'" if part.isdigit() else part for part in parts]
|
|
key = ".".join(new_parts)
|
|
|
|
jsonpath_expr = parse(key)
|
|
matches = jsonpath_expr.find(workflow_json)
|
|
|
|
if matches:
|
|
parent = matches[0].context.value
|
|
|
|
if isinstance(parent, dict):
|
|
last_part = new_parts[-1].strip("'")
|
|
parent[last_part] = value
|
|
|
|
elif isinstance(parent, list):
|
|
index = int(new_parts[-1].strip("'"))
|
|
if 0 <= index < len(parent):
|
|
parent[index] = value
|
|
else:
|
|
print(f"Index {index} is out of bounds for key '{key}'")
|
|
else:
|
|
print(f"No mapping found for key '{key}'")
|
|
|
|
return workflow_json
|
|
|
|
|
|
if __name__ == "__main__":
|
|
os.setpgrp()
|
|
try:
|
|
if not args.skip_comfy_startup:
|
|
comfyui_process = subprocess.Popen(["../venv/bin/python", "main.py", "--highvram", "--disable-metadata"], cwd=COMFY_ROOT_DIRECTORY)
|
|
main()
|
|
finally:
|
|
if not args.skip_comfy_startup:
|
|
os.killpg(0, signal.SIGKILL) # kill all processes in my group
|