mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
490 lines
20 KiB
Python
490 lines
20 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", 60))
|
|
|
|
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")
|
|
|
|
|
|
MAIN_IPA_WORKFLOW = os.environ.get("MAIN_IPA_WORKFLOW", "3-06-2024/yae_vid2vid_main_29-06_API.json")
|
|
FACE_DETAILER_WORKFLOW = os.environ.get("FACE_DETAILER_WORKFLOW", "3-06-2024/yae_vid2vid_FaceDetailer_1-07_API.json")
|
|
UPSCALER_WORKFLOW = os.environ.get("UPSCALER_WORKFLOW", "3-06-2024/yae_vid2vid_Upscale_28-06_API.json")
|
|
CINEMATIC_WORKFLOW = os.environ.get("CINEMATIC_WORKFLOW", "3-06-2024/yae_vid2vid_better_main+upscaler_04-07_API.json")
|
|
|
|
class PipelineType(IntEnum):
|
|
BASE = auto()
|
|
IPA = auto()
|
|
FACE_DETAILER = auto()
|
|
UPSCALER = auto()
|
|
CINEMATIC = auto()
|
|
|
|
|
|
style_to_filename = {
|
|
"anime_2_5d": "1_2.5d_anime_model.json",
|
|
"anime_2d_flat": "2_2d_flat_anime_model.json",
|
|
"cartoon_3d": "3_3d_cartoon_style.json",
|
|
"comic_book": "4_comic_book_model.json",
|
|
"anime_ghibli": "5_ghibli_anime_model.json",
|
|
"ink_punk": "6_ink_punk.json",
|
|
"ink_splash": "7_ink_splash.json",
|
|
"ink_bw_style": "8_ink_w_and_b_style.json",
|
|
"jojo_style": "9_jojo_style.json",
|
|
"paper_origami": "10_paper_origami.json",
|
|
"pixel_art": "11_pixel_art.json",
|
|
"pop_art": "12_pop_art.json",
|
|
"realistic_1": "13_realistic_1.json",
|
|
"realistic_2": "14_realistic_2.json",
|
|
"anime_retro_neon": "15_retro_neon_anime_90.json",
|
|
"anime_standard": "16_standard_anime_model.json",
|
|
"hr_giger": "17_hr_giger.json",
|
|
"simpsons": "18_simpsons.json",
|
|
"carnage": "19_carnage.json",
|
|
"pastel_cute_anime": "20_pastel_cute_anime.json",
|
|
"bloom_lighting": "21_bloom_lighting.json",
|
|
"25d_horror": "22_25D_Horror.json",
|
|
"creepy": "23_creepy.json",
|
|
"creepy_vhs": "24_creepy_vhs.json",
|
|
"trail_cam_footage": "25_trail_cam_footage.json",
|
|
"old_black_white_movie": "26_old_black_white_movie.json",
|
|
"horror_noir_black_white": "27_horror_noir_black_white.json",
|
|
"techno_noir_black_white": "28_techno_noir_black_white.json",
|
|
"black_white_20s": "29_black_white_20s.json",
|
|
"cyberpunk_anime": "30_cyberpunk_anime.json",
|
|
"dragonball": "31_dragonball.json",
|
|
"realistic_matrix": "32_realistic_matrix.json",
|
|
"realistic_cyberpunk": "33_realistic_cyberpunk.json"
|
|
}
|
|
|
|
|
|
def ensure_pipeline_input_present():
|
|
left = Path(COMFY_ROOT_DIRECTORY) / "input/input.mp4"
|
|
right = Path(COMFY_ROOT_DIRECTORY) / "ComfyUI" / "input/input.mp4"
|
|
shutil.copy(left, right)
|
|
|
|
|
|
def convert_pipline_output_to_input(left):
|
|
right = Path(COMFY_ROOT_DIRECTORY) / "input/input.mp4"
|
|
shutil.copy(left, right)
|
|
|
|
|
|
def restore_last_pipeline_output(left):
|
|
print("Restoring last pipeline output to input from: " + str(left))
|
|
right = Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid/SparseUpscaleInterp_00001.mp4"
|
|
shutil.copy(left, right)
|
|
|
|
|
|
def parse_args():
|
|
parser = argparse.ArgumentParser(description='Run Comfy')
|
|
parser.add_argument('--prompt', type=str, help='location of the prompt json', required=False)
|
|
parser.add_argument('--style', type=str, help='style name', required=False)
|
|
parser.add_argument('--global_ipa_image_filename', type=str, help='Global IPA image (optional)', required=False)
|
|
parser.add_argument('--global_ipa_strength', type=float, help='Global IPA strength (optional)', required=False, default=1.0)
|
|
parser.add_argument('--positive_prompt_filename', type=str, help='positive prompt', required=False)
|
|
parser.add_argument('--negative_prompt_filename', type=str, help='negative prompt', required=False)
|
|
parser.add_argument('--travel_prompt_filename', type=str, help='travel prompt', required=False)
|
|
parser.add_argument('--pause-between-steps', type=int, help='pause between steps', required=False, default=1)
|
|
parser.add_argument('--face-detailer-enabled', help='face detailer enabled', required=False, default=False, action=argparse.BooleanOptionalAction)
|
|
parser.add_argument('--upscaler-enabled', help='upscale enabled', required=False, default=False, action=argparse.BooleanOptionalAction)
|
|
parser.add_argument('--lipsync-enabled', help='lipsync enabled', required=False, default=False, action=argparse.BooleanOptionalAction)
|
|
parser.add_argument('--disable-lcm', help='disable lcm (only applies to main workflow)', required=False, default=False, action=argparse.BooleanOptionalAction)
|
|
parser.add_argument('--strength', type=float, help='strength', required=False, default=1.0)
|
|
parser.add_argument('--frame_skip', type=int, help='frame skipping', required=False, default=None)
|
|
parser.add_argument('--enable-cinematic', help='enable cinematic', required=False, default=False, action=argparse.BooleanOptionalAction)
|
|
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
|
|
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")
|
|
|
|
def interrupt_comfy_queue(server_address):
|
|
requests.post(f"http://{server_address}/interrupt")
|
|
print("Queue interrupted successfully.")
|
|
|
|
def invoke_facefusion(input_audio, input_video, output_video):
|
|
# invoke facefusion
|
|
command = [
|
|
'/app/facefusion-runner.sh',
|
|
'--input_audio', input_audio,
|
|
'--input_video', input_video,
|
|
'--output', output_video
|
|
]
|
|
subprocess.check_output(command)
|
|
print("Facefusion completed successfully.")
|
|
|
|
def main():
|
|
# interrupt_comfy_queue(SERVER_IP)
|
|
ensure_server_is_healthy(SERVER_IP)
|
|
|
|
pipeline_type = PipelineType.IPA
|
|
|
|
output_index = 0
|
|
file_path = args.prompt
|
|
prompt = None
|
|
if args.style is not None:
|
|
print("Generating prompt json")
|
|
style = args.style
|
|
positive_prompt_filename = args.positive_prompt_filename
|
|
negative_prompt_filename = args.negative_prompt_filename
|
|
travel_prompt_filename = args.travel_prompt_filename
|
|
positive_prompt = None
|
|
negative_prompt = None
|
|
travel_prompt = None
|
|
workflow_filename = None
|
|
enable_lipsync = args.lipsync_enabled
|
|
disable_lcm = args.disable_lcm
|
|
cinematic_workflow_enabled = args.enable_cinematic
|
|
denoise_first_pass = args.strength
|
|
if pipeline_type == PipelineType.IPA:
|
|
workflow_filename = MAIN_IPA_WORKFLOW
|
|
if positive_prompt_filename:
|
|
positive_prompt = open(positive_prompt_filename).read()
|
|
if negative_prompt_filename:
|
|
negative_prompt = open(negative_prompt_filename).read()
|
|
if travel_prompt_filename:
|
|
travel_prompt = open(travel_prompt_filename).read()
|
|
if cinematic_workflow_enabled:
|
|
pipeline_type = PipelineType.CINEMATIC
|
|
workflow_filename = CINEMATIC_WORKFLOW
|
|
|
|
time_before = time.perf_counter()
|
|
|
|
validate_style_name(style)
|
|
prompt = generate_prompt_for_style(
|
|
style,
|
|
positive_prompt,
|
|
negative_prompt,
|
|
travel_prompt,
|
|
workflow_filename,
|
|
denoise_first_pass=denoise_first_pass,
|
|
enable_lipsync=enable_lipsync,
|
|
disable_lcm=disable_lcm,
|
|
global_ipa_image_filename=args.global_ipa_image_filename,
|
|
global_ipa_strength=args.global_ipa_strength,
|
|
frame_skip=args.frame_skip
|
|
)
|
|
print("running pipeline" + str(pipeline_type))
|
|
enqueue_prompt_and_wait(prompt)
|
|
output_index = output_index + 1
|
|
left = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
|
if args.upscaler_enabled:
|
|
if pipeline_type is not PipelineType.CINEMATIC:
|
|
print("running pipeline" + str(pipeline_type))
|
|
convert_pipline_output_to_input(left)
|
|
workflow_filename = UPSCALER_WORKFLOW
|
|
prompt = generate_prompt_for_style(
|
|
style,
|
|
positive_prompt,
|
|
negative_prompt,
|
|
travel_prompt,
|
|
workflow_filename,
|
|
PipelineType.UPSCALER,
|
|
denoise_first_pass=denoise_first_pass,
|
|
enable_lipsync=enable_lipsync,
|
|
global_ipa_image_filename=args.global_ipa_image_filename,
|
|
global_ipa_strength=args.global_ipa_strength,
|
|
frame_skip=args.frame_skip
|
|
)
|
|
enqueue_prompt_and_wait(prompt)
|
|
output_index = output_index + 1
|
|
left = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
|
if args.face_detailer_enabled:
|
|
print("running pipeline" + str(pipeline_type))
|
|
convert_pipline_output_to_input(left)
|
|
workflow_filename = FACE_DETAILER_WORKFLOW
|
|
prompt = generate_prompt_for_style(
|
|
style,
|
|
positive_prompt,
|
|
negative_prompt,
|
|
travel_prompt,
|
|
workflow_filename,
|
|
PipelineType.FACE_DETAILER,
|
|
denoise_first_pass=denoise_first_pass,
|
|
enable_lipsync=enable_lipsync,
|
|
global_ipa_image_filename=args.global_ipa_image_filename,
|
|
global_ipa_strength=args.global_ipa_strength,
|
|
frame_skip=args.frame_skip
|
|
)
|
|
enqueue_prompt_and_wait(prompt)
|
|
output_index = output_index + 1
|
|
left = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
|
if enable_lipsync:
|
|
# Ugh
|
|
input_audio = Path(COMFY_ROOT_DIRECTORY) / "input" / "trimmed.wav"
|
|
input_video = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
|
output_index = output_index + 1
|
|
output_video = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
|
try:
|
|
invoke_facefusion(input_audio, input_video, output_video)
|
|
except Exception as e:
|
|
print("Facefusion failed!")
|
|
print(e)
|
|
# check if output_video file exists
|
|
if not output_video.exists():
|
|
print("Facefusion failed!")
|
|
else:
|
|
left = output_video
|
|
restore_last_pipeline_output(left)
|
|
time_after = time.perf_counter()
|
|
print(f"Time taken: {time_after - time_before:.2f} seconds")
|
|
else:
|
|
print(f'Loading prompt json from path: {file_path}')
|
|
prompt_workflow = json.load(open(file_path))
|
|
prompt = {"prompt": prompt_workflow}
|
|
enqueue_prompt_and_wait(prompt)
|
|
|
|
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(style_name,
|
|
positive_prompt,
|
|
negative_prompt,
|
|
travel_prompt,
|
|
workflow_filename=None,
|
|
pipeline_type = PipelineType.BASE,
|
|
denoise_first_pass=1.0,
|
|
enable_lipsync: bool = False,
|
|
disable_lcm: bool = False,
|
|
global_ipa_image_filename: Optional[str] = None,
|
|
global_ipa_strength: float = 1.0,
|
|
frame_skip: Optional[int] = None,
|
|
) -> Path:
|
|
styles_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "styles"
|
|
workflow_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "workflows"
|
|
mappings_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "mappings"
|
|
|
|
style_filename = style_to_filename[style_name]
|
|
style_json = json.load(open(styles_directory / style_filename))
|
|
|
|
workflow_json_filename = style_json["workflow_api_name"]
|
|
mapping_json_filename = "mapping_one.json"
|
|
|
|
if workflow_filename:
|
|
workflow_json_filename = workflow_filename
|
|
|
|
workflow_json = json.load(open(workflow_directory / workflow_json_filename))
|
|
mapping_json = json.load(open(mappings_directory / mapping_json_filename))
|
|
|
|
json_mods = {key: value for key, value in style_json["modifications"].items() if not key.startswith("cn_")}
|
|
|
|
jsonpath_mods = get_jsonpath_mods(
|
|
json_mods,
|
|
mapping_json,
|
|
pos_in=positive_prompt,
|
|
neg_in=negative_prompt,
|
|
travel_in=travel_prompt,
|
|
pipeline_type=pipeline_type,
|
|
denoise_first_pass=denoise_first_pass,
|
|
enable_lipsync=enable_lipsync,
|
|
disable_lcm=disable_lcm,
|
|
global_ipa_image_filename=global_ipa_image_filename,
|
|
global_ipa_strength=global_ipa_strength,
|
|
frame_skip=frame_skip,
|
|
)
|
|
|
|
print(*jsonpath_mods.items(), sep="\n")
|
|
|
|
json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json)
|
|
|
|
p = {"prompt": json_prompt}
|
|
return p
|
|
|
|
|
|
def get_jsonpath_mods(style_mods: Dict[str, Any],
|
|
mapping_json: Dict[str, Any],
|
|
pos_in: Optional[str] = None,
|
|
neg_in: Optional[str] = None,
|
|
travel_in: Optional[str] = None,
|
|
pipeline_type = PipelineType.BASE,
|
|
denoise_first_pass=1.0,
|
|
enable_lipsync: bool = False,
|
|
disable_lcm: bool = False,
|
|
global_ipa_image_filename: Optional[str] = None,
|
|
global_ipa_strength: float = 1.0,
|
|
frame_skip: Optional[int] = None,
|
|
) -> Dict[str, Any]:
|
|
modifications = {}
|
|
new_mod_json = style_mods.copy()
|
|
# Process "loras" differently
|
|
loras = new_mod_json.get("loras", [])
|
|
if len(loras) > 8:
|
|
raise ValueError("Too many loras, max is 8")
|
|
|
|
for index, lora in enumerate(loras):
|
|
if "name" in lora and "strength" in lora:
|
|
lora_strength = lora["strength"]
|
|
if pipeline_type == PipelineType.FACE_DETAILER:
|
|
lora_strength = 0.5
|
|
new_mod_json[f"lora_{index + 1}_strength"] = lora_strength
|
|
new_mod_json[f"lora_{index + 1}_name"] = lora["name"]
|
|
|
|
# Exclude "loras" from further processing
|
|
new_mod_json.pop("loras", None)
|
|
|
|
# Handling positive and negative prompts
|
|
if pos_in:
|
|
new_mod_json["positive_prompt"] = f"{pos_in}, {new_mod_json.get('positive_prompt', '')},"
|
|
if neg_in:
|
|
new_mod_json["negative_prompt"] = f"{neg_in}, {new_mod_json.get('negative_prompt', '')},"
|
|
|
|
if pipeline_type in [PipelineType.IPA, PipelineType.BASE]:
|
|
new_mod_json["denoise_first_pass"] = denoise_first_pass
|
|
|
|
if enable_lipsync:
|
|
new_mod_json["cn_lips_strength"] = 1.2
|
|
|
|
if frame_skip:
|
|
new_mod_json["every_nth_frame"] = frame_skip
|
|
|
|
for key, value in new_mod_json.items():
|
|
mapping_key = f"$.{key}"
|
|
jsonpath_expr = parse(mapping_key)
|
|
|
|
# Finding matches using jsonpath_ng
|
|
matches = [match.value for match in jsonpath_expr.find(mapping_json)]
|
|
if matches:
|
|
mapping_value = matches[0]
|
|
modifications[mapping_value] = value
|
|
else:
|
|
print(f"No mapping found for key '{key}'")
|
|
|
|
# TODO(bt,2024-05-29): Move this to the mappings so we don't have to maintain the jsonpath
|
|
# This is only relevant for the main workflow (not face fixer or upscaler)
|
|
if disable_lcm:
|
|
modifications["$.536.inputs.boolean_number"] = 0
|
|
|
|
# TODO(bt,2024-06-22): Move this to the mappings so we don't have to maintain the jsonpath
|
|
if global_ipa_image_filename:
|
|
modifications["$.3723.inputs.image_path"] = global_ipa_image_filename
|
|
modifications["$.900.inputs.value"] = global_ipa_strength
|
|
|
|
# TODO(bt,2024-07-06): Move this to the mappings so we don't have to maintain the json path
|
|
if travel_in:
|
|
modifications["$.509.inputs.text"] = travel_in
|
|
|
|
return modifications
|
|
|
|
|
|
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
|
|
|
|
|
|
def validate_style_name(style):
|
|
if style not in style_to_filename:
|
|
print(f"Invalid style: {args.style}")
|
|
print(f"Valid styles: {', '.join(style_to_filename.keys())}")
|
|
exit(1)
|
|
|
|
if __name__ == "__main__":
|
|
os.setpgrp()
|
|
try:
|
|
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
|