|
|
|
@@ -13,7 +13,7 @@ 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", 180))
|
|
|
|
|
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"
|
|
|
|
@@ -25,10 +25,7 @@ COMFY_ROOT_DIRECTORY = os.environ.get("COMFY_ROOT_DIRECTORY",
|
|
|
|
|
"/app/ComfyUI")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
BASE_LIVE_PORTRAIT_WORKFLOW = os.environ.get("BASE_LIVE_PORTRAIT_WORKFLOW", "17-07-2024/yae_LivingPortrait_17-07_API.json")
|
|
|
|
|
|
|
|
|
|
# V2V_WORKFLOWS_DIRECTORY = "/app/ComfyUI/exp"
|
|
|
|
|
# BASE_LIVE_PORTRAIT_WORKFLOW = "yae_LivingPortrait_17-07_API.json"
|
|
|
|
|
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()
|
|
|
|
@@ -56,12 +53,18 @@ def ensure_pipeline_input_present(source_path, input_name=None):
|
|
|
|
|
right = Path(COMFY_ROOT_DIRECTORY) / "input" / right_filename
|
|
|
|
|
shutil.copy(left, right)
|
|
|
|
|
|
|
|
|
|
def primary_output_directory():
|
|
|
|
|
return Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid"
|
|
|
|
|
|
|
|
|
|
def ensure_output_present_at(new_filename):
|
|
|
|
|
directory = 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 = directory / "LivePortrait_00001.mp4"
|
|
|
|
|
output_with_audio = directory / "LivePortrait_00001-audio.mp4"
|
|
|
|
|
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)
|
|
|
|
@@ -76,67 +79,130 @@ def parse_args():
|
|
|
|
|
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 ensure_server_is_healthy(server_address):
|
|
|
|
|
# check the health of the server (5 second timeout)
|
|
|
|
|
response = requests.get(f"http://{server_address}/queue", timeout=5)
|
|
|
|
|
if response.status_code == 200:
|
|
|
|
|
print("Server is healthy")
|
|
|
|
|
else:
|
|
|
|
|
print("Server is unhealthy")
|
|
|
|
|
print(f"Status code: {response.status_code}")
|
|
|
|
|
print(f"Response: {response.text}")
|
|
|
|
|
|
|
|
|
|
# make sure there are no jobs already
|
|
|
|
|
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()
|
|
|
|
|
if len(queue_json.get("queue_pending", [])) > 0:
|
|
|
|
|
print("There are jobs in the queue already. Exiting...")
|
|
|
|
|
exit(1)
|
|
|
|
|
else:
|
|
|
|
|
print("Queue is empty")
|
|
|
|
|
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():
|
|
|
|
|
ensure_server_is_healthy(SERVER_IP)
|
|
|
|
|
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
|
|
|
|
|
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")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
time_before = time.perf_counter()
|
|
|
|
|
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,
|
|
|
|
|
)
|
|
|
|
|
print("running pipeline" + str(pipeline_type))
|
|
|
|
|
enqueue_prompt_and_wait(prompt)
|
|
|
|
|
for _ in range(outer_retries):
|
|
|
|
|
ensure_server_is_healthy(SERVER_IP, restart_trigger_file)
|
|
|
|
|
|
|
|
|
|
ensure_output_present_at(args.output_filename)
|
|
|
|
|
time_after = time.perf_counter()
|
|
|
|
|
print(f"Time taken: {time_after - time_before:.2f} seconds")
|
|
|
|
|
ensure_input_directory_empty()
|
|
|
|
|
ensure_output_directory_empty()
|
|
|
|
|
|
|
|
|
|
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")
|
|
|
|
|
|
|
|
|
|
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')
|
|
|
|
@@ -242,21 +308,7 @@ if __name__ == "__main__":
|
|
|
|
|
try:
|
|
|
|
|
if not args.skip_comfy_startup:
|
|
|
|
|
comfyui_process = subprocess.Popen(["../venv/bin/python", "main.py", "--highvram", "--disable-metadata"], cwd=COMFY_ROOT_DIRECTORY)
|
|
|
|
|
start_time = time.time()
|
|
|
|
|
started = False
|
|
|
|
|
while time.time() - start_time < SERVER_START_TIMEOUT:
|
|
|
|
|
try:
|
|
|
|
|
response = requests.get(PROMPT_ENDPOINT, timeout=5)
|
|
|
|
|
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)
|
|
|
|
|
main()
|
|
|
|
|
finally:
|
|
|
|
|
os.killpg(0, signal.SIGKILL) # kill all processes in my group
|
|
|
|
|
if not args.skip_comfy_startup:
|
|
|
|
|
os.killpg(0, signal.SIGKILL) # kill all processes in my group
|
|
|
|
|