mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
comfy: write preview frames where rust can find them
This commit is contained in:
@@ -11,11 +11,16 @@ from enum import IntEnum, auto
|
||||
import os
|
||||
import subprocess
|
||||
import signal
|
||||
import urllib
|
||||
import websockets
|
||||
import asyncio
|
||||
import uuid
|
||||
client_id = str(uuid.uuid4())
|
||||
|
||||
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))
|
||||
TIMEOUT_SECONDS = int(os.environ.get("TIMEOUT_SECONDS", 180))
|
||||
PROMPT_ENDPOINT = f"http://{SERVER_IP}/prompt"
|
||||
V2V_WORKFLOWS_DIRECTORY = os.environ.get("V2V_WORKFLOWS_DIRECTORY",
|
||||
"/workflow_configs")
|
||||
@@ -120,6 +125,7 @@ def ensure_server_is_healthy(server_address, restart_trigger_file=None, retries=
|
||||
|
||||
for _ in range(retries):
|
||||
trigger_comfy_restart(restart_trigger_file)
|
||||
time.sleep(10)
|
||||
wait_for_comfy_startup(server_address)
|
||||
number_of_queued_jobs = get_queued_jobs_count(server_address)
|
||||
if queued_count_is_zero(server_address):
|
||||
@@ -153,7 +159,7 @@ def wait_for_comfy_startup(server_address, timeout=SERVER_START_TIMEOUT):
|
||||
exit(1)
|
||||
|
||||
|
||||
def main():
|
||||
async def main():
|
||||
outer_retries = 2
|
||||
restart_trigger_file = args.server_restart_trigger_file
|
||||
|
||||
@@ -163,7 +169,7 @@ def main():
|
||||
input_driver_filename = args.driver_media_filename
|
||||
input_is_image = args.input_is_image
|
||||
|
||||
for _ in range(outer_retries):
|
||||
for attempt_idx in range(outer_retries):
|
||||
ensure_server_is_healthy(SERVER_IP, restart_trigger_file)
|
||||
|
||||
ensure_input_directory_empty()
|
||||
@@ -188,7 +194,7 @@ def main():
|
||||
)
|
||||
time_before = time.perf_counter()
|
||||
print("running pipeline" + str(pipeline_type))
|
||||
enqueue_prompt_and_wait(prompt)
|
||||
await 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():
|
||||
@@ -200,11 +206,34 @@ def main():
|
||||
if restart_trigger_file is None:
|
||||
print("No restart trigger file provided. Exiting.")
|
||||
exit(1)
|
||||
trigger_comfy_restart(restart_trigger_file)
|
||||
if attempt_idx < outer_retries - 1:
|
||||
trigger_comfy_restart(restart_trigger_file)
|
||||
time.sleep(10)
|
||||
# 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):
|
||||
def classify_error_message(message_data):
|
||||
if message_data['node_type'] == 'LivePortraitCropper':
|
||||
print("Error in LivePortraitCropper")
|
||||
|
||||
async def smart_wait(ws, prompt_id):
|
||||
while True:
|
||||
out = await ws.recv()
|
||||
if isinstance(out, str):
|
||||
message = json.loads(out)
|
||||
if message['type'] != 'progress':
|
||||
print(message)
|
||||
if message['type'] == 'execution_error':
|
||||
classify_error_message(message['data'])
|
||||
if message['type'] == 'executing':
|
||||
data = message['data']
|
||||
if data['node'] is None and data['prompt_id'] == prompt_id:
|
||||
break #Execution is done
|
||||
else:
|
||||
continue #previews are binary data
|
||||
|
||||
|
||||
async def enqueue_prompt_and_wait(prompt):
|
||||
data = json.dumps(prompt).encode('utf-8')
|
||||
# Send HTTP post request
|
||||
response = requests.post(PROMPT_ENDPOINT, data=data)
|
||||
@@ -215,29 +244,20 @@ def enqueue_prompt_and_wait(prompt):
|
||||
print(f"Status code: {response.status_code}")
|
||||
print(f"Response: {response.text}")
|
||||
exit(1)
|
||||
|
||||
prompt_id = response.json()["prompt_id"]
|
||||
# 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)
|
||||
execution_start_time = time.perf_counter()
|
||||
|
||||
async with websockets.connect("ws://{}/ws?clientId={}".format(SERVER_IP, client_id)) as ws:
|
||||
await smart_wait(ws, prompt_id)
|
||||
current_time = time.perf_counter()
|
||||
execution_time = current_time - execution_start_time
|
||||
print('request executed in {:.2f} seconds'.format(execution_time))
|
||||
|
||||
def get_history(prompt_id):
|
||||
with urllib.request.urlopen("http://{}/history/{}".format(SERVER_IP, prompt_id)) as response:
|
||||
return json.loads(response.read())
|
||||
|
||||
def generate_prompt_for_style(
|
||||
workflow_filename: str,
|
||||
@@ -270,7 +290,7 @@ def generate_prompt_for_style(
|
||||
|
||||
json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json)
|
||||
|
||||
p = {"prompt": json_prompt}
|
||||
p = {"prompt": json_prompt, "client_id": client_id}
|
||||
return p
|
||||
|
||||
def apply_jsonpath_mods(jsonpath_mods: Dict[str, Any], workflow_json: Dict[str, Any]) -> Dict[str, Any]:
|
||||
@@ -307,7 +327,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)
|
||||
main()
|
||||
asyncio.run(main())
|
||||
finally:
|
||||
if not args.skip_comfy_startup:
|
||||
os.killpg(0, signal.SIGKILL) # kill all processes in my group
|
||||
|
||||
@@ -11,6 +11,15 @@ from enum import IntEnum, auto
|
||||
import os
|
||||
import subprocess
|
||||
import signal
|
||||
import websockets
|
||||
import httpx
|
||||
import asyncio
|
||||
import urllib
|
||||
import uuid
|
||||
import re
|
||||
import imageio.v3 as iio
|
||||
|
||||
client_id = str(uuid.uuid4())
|
||||
|
||||
SERVER_IP = os.environ.get("SERVER_IP", "127.0.0.1:8188")
|
||||
SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 60))
|
||||
@@ -34,6 +43,7 @@ class PipelineType(IntEnum):
|
||||
FACE_DETAILER = auto()
|
||||
UPSCALER = auto()
|
||||
CINEMATIC = auto()
|
||||
PREVIEW = auto()
|
||||
|
||||
|
||||
style_to_filename = {
|
||||
@@ -79,18 +89,82 @@ def ensure_pipeline_input_present():
|
||||
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 copy_frames_to_directory(stage, frames_dir):
|
||||
frames = Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid" / stage / "Frames"
|
||||
frame_idx_matcher = re.compile('.*Frame_(\d+)*')
|
||||
for file in frames.iterdir():
|
||||
# output/vid2vid/FirstPass/Frame_00002_.png
|
||||
# output/vid2vid/SecondPass/Frames/Frame_0006.jpg
|
||||
# remove trailing prefiix and Frame_ prefix
|
||||
file_name = file.name
|
||||
frame_idx = frame_idx_matcher.match(file_name).group(1)
|
||||
target_file_name = f"{frame_idx.zfill(5)}{file.suffix}"
|
||||
shutil.copy(file, frames_dir / target_file_name)
|
||||
print(f"Copying {file} to {frames_dir / target_file_name}")
|
||||
|
||||
def generated_frames_count(frames_dir):
|
||||
return len(list(frames_dir.glob("*.jpg")))
|
||||
|
||||
# ffmpeg -i input.mp4 -qscale:v 2 output_%03d.jpg
|
||||
async def extract_frames(input_video, output_directory):
|
||||
if not output_directory.exists():
|
||||
output_directory.mkdir(parents=True)
|
||||
for idx, frame in enumerate(iio.imiter(input_video)):
|
||||
frame_idx = idx + 1
|
||||
iio.imwrite(f"{output_directory}/Frame_{frame_idx:05d}.jpg", frame)
|
||||
# command = [
|
||||
# 'ffmpeg',
|
||||
# '-i', str(input_video),
|
||||
# '-qscale:v', '2',
|
||||
# str(f'{output_directory}/Frame_%03d.jpg')
|
||||
# ]
|
||||
# await cmd_runner(" ".join(command))
|
||||
print("Frames extracted successfully.")
|
||||
|
||||
def pad_frames(frames_dir, expected_count):
|
||||
# Copy last frames over till we reach expected_count
|
||||
# find first frame that exists
|
||||
last_frame = None
|
||||
last_frame_idx = 0
|
||||
for i in range(expected_count, 0, -1):
|
||||
frame = frames_dir / f"{i:05d}.jpg"
|
||||
if frame.exists():
|
||||
last_frame = frame
|
||||
last_frame_idx = i
|
||||
break
|
||||
for idx in range(last_frame_idx + 1, expected_count + 1):
|
||||
print(f"Copying padding frame {last_frame} to Frame_{idx:05d}.jpg")
|
||||
shutil.copy(last_frame, frames_dir / f"{idx:05d}.jpg")
|
||||
|
||||
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 ensure_frames_dir_present(stage):
|
||||
# output/vid2vid/FirstPass/Frames
|
||||
frames_dir = Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid" / stage / "Frames"
|
||||
if not frames_dir.exists():
|
||||
frames_dir.mkdir(parents=True)
|
||||
|
||||
def ensure_frames_dir_empty(stage):
|
||||
frames_dir = Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid" / stage / "Frames"
|
||||
if frames_dir.exists():
|
||||
for file in frames_dir.iterdir():
|
||||
try:
|
||||
file.unlink()
|
||||
except Exception as e:
|
||||
print(f"Error while deleting {file}")
|
||||
print(e)
|
||||
exit(1)
|
||||
|
||||
def ensure_output_directory_empty():
|
||||
shutil.rmtree(Path(COMFY_ROOT_DIRECTORY) / "output")
|
||||
(Path(COMFY_ROOT_DIRECTORY) / "output").mkdir(parents=True)
|
||||
|
||||
def parse_args():
|
||||
|
||||
@@ -113,9 +187,10 @@ def parse_args():
|
||||
parser.add_argument('--depth_video_filename', type=str, help='path of mp4 depth from engine for preprocessing', required=False)
|
||||
parser.add_argument('--outline_video_filename', type=str, help='path of mp4 outline from engine for preprocessing', required=False)
|
||||
parser.add_argument('--normals_video_filename', type=str, help='path of mp4 normals from engine for preprocessing', required=False)
|
||||
|
||||
parser.add_argument('--generate-previews', help='generate previews', required=False, default=False, action=argparse.BooleanOptionalAction)
|
||||
parser.add_argument('--preview-frames-directory', type=str, help='previews directory', required=False)
|
||||
parser.add_argument('--enable-cinematic', help='enable cinematic', required=False, default=False, action=argparse.BooleanOptionalAction)
|
||||
|
||||
parser.add_argument('--skip-comfy-startup', action='store_true', help='controls if this script tries to start comfy', default=False)
|
||||
parser.add_argument('--server-restart-trigger-file', type=str, help='server restart trigger file', required=False)
|
||||
return parser.parse_args()
|
||||
|
||||
@@ -182,9 +257,20 @@ def invoke_facefusion(input_audio, input_video, output_video):
|
||||
subprocess.check_output(command)
|
||||
print("Facefusion completed successfully.")
|
||||
|
||||
def main():
|
||||
async def main():
|
||||
if not args.skip_comfy_startup:
|
||||
print("Starting comfy server")
|
||||
try:
|
||||
comfy_process = subprocess.Popen(["../venv/bin/python", "main.py", "--highvram", "--disable-metadata"], cwd=COMFY_ROOT_DIRECTORY)
|
||||
except Exception as e:
|
||||
print("Error while starting comfy server")
|
||||
print(e)
|
||||
exit(1)
|
||||
await asyncio.sleep(15)
|
||||
else:
|
||||
print("Skipping comfy server startup")
|
||||
# interrupt_comfy_queue(SERVER_IP)
|
||||
ensure_server_is_healthy(SERVER_IP)
|
||||
wait_for_comfy_startup(SERVER_IP)
|
||||
|
||||
pipeline_type = PipelineType.IPA
|
||||
|
||||
@@ -205,22 +291,57 @@ def main():
|
||||
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
|
||||
# defaults to none
|
||||
depth_video_filename = args.depth_video_filename
|
||||
outline_video_filename = args.outline_video_filename
|
||||
normals_video_filename = args.normals_video_filename
|
||||
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()
|
||||
ensure_output_directory_empty()
|
||||
if args.generate_previews:
|
||||
previews_directory = args.preview_frames_directory
|
||||
if not previews_directory:
|
||||
print("Previews directory not set")
|
||||
exit(1)
|
||||
ensure_frames_dir_present("FirstPass")
|
||||
if not (Path(previews_directory) / "first_pass" ).exists():
|
||||
(Path(previews_directory) / "first_pass").mkdir(parents=True)
|
||||
ensure_frames_dir_empty("FirstPass")
|
||||
workflow_filename = PREVIEW_IPA_WORKFLOW
|
||||
pipeline_type = PipelineType.PREVIEW
|
||||
prompt = generate_prompt_for_style(
|
||||
style,
|
||||
positive_prompt,
|
||||
negative_prompt,
|
||||
travel_prompt,
|
||||
workflow_filename,
|
||||
PipelineType.PREVIEW,
|
||||
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,
|
||||
depth_video_filename=depth_video_filename,
|
||||
outline_video_filename=outline_video_filename,
|
||||
normals_video_filename=normals_video_filename
|
||||
)
|
||||
print("running pipeline" + str(pipeline_type))
|
||||
await enqueue_prompt_and_wait(prompt)
|
||||
copy_frames_to_directory("FirstPass", Path(previews_directory) / "first_pass")
|
||||
pipeline_type = PipelineType.IPA
|
||||
workflow_filename = MAIN_IPA_WORKFLOW
|
||||
|
||||
ensure_output_directory_empty()
|
||||
if pipeline_type == PipelineType.IPA:
|
||||
workflow_filename = MAIN_IPA_WORKFLOW
|
||||
if cinematic_workflow_enabled:
|
||||
pipeline_type = PipelineType.CINEMATIC
|
||||
workflow_filename = CINEMATIC_WORKFLOW
|
||||
|
||||
# defaults to none
|
||||
depth_video_filename = args.depth_video_filename
|
||||
outline_video_filename = args.outline_video_filename
|
||||
normals_video_filename = args.normals_video_filename
|
||||
|
||||
time_before = time.perf_counter()
|
||||
|
||||
@@ -244,11 +365,17 @@ def main():
|
||||
)
|
||||
|
||||
print("running pipeline" + str(pipeline_type))
|
||||
enqueue_prompt_and_wait(prompt)
|
||||
await enqueue_prompt_and_wait(prompt)
|
||||
output_index = output_index + 1
|
||||
left = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
||||
if left.exists():
|
||||
print("Pipeline completed successfully to output file: " + str(left))
|
||||
else:
|
||||
print("Pipeline failed to produce output")
|
||||
exit(1)
|
||||
if args.upscaler_enabled:
|
||||
if pipeline_type is not PipelineType.CINEMATIC:
|
||||
pipeline_type = PipelineType.UPSCALER
|
||||
print("running pipeline" + str(pipeline_type))
|
||||
convert_pipline_output_to_input(left)
|
||||
workflow_filename = UPSCALER_WORKFLOW
|
||||
@@ -258,17 +385,33 @@ def main():
|
||||
negative_prompt,
|
||||
travel_prompt,
|
||||
workflow_filename,
|
||||
PipelineType.UPSCALER,
|
||||
pipeline_type,
|
||||
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)
|
||||
if args.generate_previews:
|
||||
previews_directory = args.preview_frames_directory
|
||||
if not previews_directory:
|
||||
print("Previews directory not set")
|
||||
exit(1)
|
||||
if not (Path(previews_directory) / "second_pass" ).exists():
|
||||
(Path(previews_directory) / "second_pass").mkdir(parents=True)
|
||||
ensure_frames_dir_empty("SecondPass")
|
||||
await enqueue_prompt_and_wait(prompt)
|
||||
if args.generate_previews:
|
||||
copy_frames_to_directory("SecondPass", Path(previews_directory) / "second_pass")
|
||||
output_index = output_index + 1
|
||||
left = Path(COMFY_ROOT_DIRECTORY) / "output" / f"vid2vid/SparseUpscaleInterp_0000{output_index}.mp4"
|
||||
if left.exists():
|
||||
print("Pipeline completed successfully to output file: " + str(left))
|
||||
else:
|
||||
print("Pipeline failed to produce output")
|
||||
exit(1)
|
||||
if args.face_detailer_enabled:
|
||||
pipeline_type = PipelineType.FACE_DETAILER
|
||||
print("running pipeline" + str(pipeline_type))
|
||||
convert_pipline_output_to_input(left)
|
||||
workflow_filename = FACE_DETAILER_WORKFLOW
|
||||
@@ -278,14 +421,14 @@ def main():
|
||||
negative_prompt,
|
||||
travel_prompt,
|
||||
workflow_filename,
|
||||
PipelineType.FACE_DETAILER,
|
||||
pipeline_type,
|
||||
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)
|
||||
await 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:
|
||||
@@ -305,47 +448,102 @@ def main():
|
||||
else:
|
||||
left = output_video
|
||||
restore_last_pipeline_output(left)
|
||||
await extract_frames(left, Path(COMFY_ROOT_DIRECTORY) / "output" / "vid2vid" / "FinalPass" / "Frames")
|
||||
if not (Path(previews_directory) / "final_pass" ).exists():
|
||||
(Path(previews_directory) / "final_pass").mkdir(parents=True)
|
||||
copy_frames_to_directory("FinalPass", Path(previews_directory) / "final_pass")
|
||||
first_pass_frame_count = generated_frames_count(Path(previews_directory) / "first_pass")
|
||||
second_pass_frame_count = generated_frames_count(Path(previews_directory) / "second_pass")
|
||||
final_pass_frame_count = generated_frames_count(Path(previews_directory) / "final_pass")
|
||||
|
||||
if second_pass_frame_count < first_pass_frame_count:
|
||||
pad_frames(Path(previews_directory) / "second_pass", first_pass_frame_count)
|
||||
if final_pass_frame_count < first_pass_frame_count:
|
||||
pad_frames(Path(previews_directory) / "final_pass", first_pass_frame_count)
|
||||
|
||||
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)
|
||||
prompt = {"prompt": prompt_workflow, "client_id": client_id}
|
||||
await 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)
|
||||
def get_history(prompt_id):
|
||||
with urllib.request.urlopen("http://{}/history/{}".format(SERVER_IP, prompt_id)) as response:
|
||||
return json.loads(response.read())
|
||||
|
||||
# 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
|
||||
def get_queue(server_ip):
|
||||
with urllib.request.urlopen("http://{}/queue".format(server_ip)) as response:
|
||||
return json.loads(response.read())
|
||||
|
||||
def confirm_prompt_running_or_pending(prompt_id):
|
||||
queued_response = get_queue(SERVER_IP)
|
||||
prompt_possibly_queued = False
|
||||
|
||||
for possible_state in ['queue_running', 'queue_pending']:
|
||||
for j in queued_response.get(possible_state,[]):
|
||||
try:
|
||||
if j[1] == prompt_id:
|
||||
prompt_possibly_queued = True
|
||||
break
|
||||
except IndexError:
|
||||
pass
|
||||
return prompt_possibly_queued
|
||||
|
||||
|
||||
async def smart_wait(ws, prompt_id, max_wait_time=10):
|
||||
last_prompt_update_time = time.monotonic_ns()
|
||||
while True:
|
||||
try:
|
||||
response = requests.get(PROMPT_ENDPOINT, timeout=5)
|
||||
queue_json = response.json()
|
||||
if queue_json["exec_info"]["queue_remaining"]:
|
||||
pass
|
||||
out = await ws.recv()
|
||||
last_msg_time = time.monotonic_ns()
|
||||
if isinstance(out, str):
|
||||
message = json.loads(out)
|
||||
if message.get('data', None) and message['data'].get('prompt_id', None) == prompt_id:
|
||||
last_prompt_update_time = time.monotonic_ns()
|
||||
# print(f"Prompt {prompt_id} updated at {last_prompt_update_time}")
|
||||
if (last_msg_time - last_prompt_update_time > max_wait_time * 1e9) and not confirm_prompt_running_or_pending(prompt_id):
|
||||
print(f"Prompt {prompt_id} might be stuck")
|
||||
break
|
||||
if message['type'] == 'executing':
|
||||
data = message['data']
|
||||
if data['node'] is None and data['prompt_id'] == prompt_id:
|
||||
break
|
||||
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!")
|
||||
print('received binary message')
|
||||
if last_msg_time - last_prompt_update_time > max_wait_time * 1e9:
|
||||
print(f"Prompt {prompt_id} might be stuck")
|
||||
break
|
||||
continue
|
||||
except websockets.exceptions.ConnectionClosedError:
|
||||
print('Connection closed')
|
||||
break
|
||||
print(f"Prompt {prompt_id} - finished finished")
|
||||
|
||||
async def enqueue_prompt_and_wait(prompt):
|
||||
data = json.dumps(prompt).encode('utf-8')
|
||||
|
||||
with httpx.Client() as client:
|
||||
response = client.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)
|
||||
time.sleep(1)
|
||||
prompt_id = response.json()["prompt_id"]
|
||||
# Wait for the job to finish
|
||||
print(f"Waiting for job to finish (timeout: {TIMEOUT_SECONDS} seconds)")
|
||||
execution_start_time = time.perf_counter()
|
||||
# import ipdb; ipdb.set_trace()
|
||||
# max_size = 1048576 * 10 allows us to receive a really large payload from comfy
|
||||
async with websockets.connect("ws://{}/ws?clientId={}".format(SERVER_IP, client_id), max_size=1048576 * 10) as ws:
|
||||
await smart_wait(ws, prompt_id)
|
||||
current_time = time.perf_counter()
|
||||
execution_time = current_time - execution_start_time
|
||||
print('request executed in {:.2f} seconds'.format(execution_time))
|
||||
|
||||
|
||||
|
||||
@@ -405,9 +603,25 @@ def generate_prompt_for_style(style_name,
|
||||
|
||||
json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json)
|
||||
|
||||
p = {"prompt": json_prompt}
|
||||
p = {"prompt": json_prompt, "client_id": client_id}
|
||||
return p
|
||||
|
||||
async def cmd_runner(cmd: str):
|
||||
proc = await asyncio.create_subprocess_shell(
|
||||
cmd,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE
|
||||
)
|
||||
|
||||
stdout, stderr = await proc.communicate()
|
||||
|
||||
print(f'[{cmd!r} exited with {proc.returncode}]')
|
||||
if stdout:
|
||||
print(f'[stdout]\n{stdout.decode()}')
|
||||
if stderr:
|
||||
print(f'[stderr]\n{stderr.decode()}')
|
||||
|
||||
|
||||
|
||||
def get_jsonpath_mods(style_mods: Dict[str, Any],
|
||||
mapping_json: Dict[str, Any],
|
||||
@@ -450,7 +664,7 @@ def get_jsonpath_mods(style_mods: Dict[str, Any],
|
||||
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]:
|
||||
if pipeline_type in [PipelineType.IPA, PipelineType.BASE, PipelineType.PREVIEW]:
|
||||
new_mod_json["denoise_first_pass"] = denoise_first_pass
|
||||
|
||||
if enable_lipsync:
|
||||
@@ -487,7 +701,7 @@ def get_jsonpath_mods(style_mods: Dict[str, Any],
|
||||
modifications["$.509.inputs.text"] = travel_in
|
||||
|
||||
# TODO The 3 Mappings for depth outlines and normals apparently uses an array.
|
||||
if pipeline_type in [PipelineType.IPA, PipelineType.BASE]:
|
||||
if pipeline_type in [PipelineType.IPA, PipelineType.BASE, PipelineType.PREVIEW]:
|
||||
if depth_video_filename != None and outline_video_filename != None and normals_video_filename !=None:
|
||||
modifications["$.3731.inputs.text"] = depth_video_filename
|
||||
modifications["$.4120.inputs.text"] = normals_video_filename
|
||||
@@ -533,22 +747,9 @@ def validate_style_name(style):
|
||||
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()
|
||||
asyncio.run(main())
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
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
|
||||
|
||||
@@ -66,7 +66,7 @@
|
||||
},
|
||||
{
|
||||
"folder": "ComfyUI/custom_nodes/ComfyUI_essentials",
|
||||
"commit": "0283ce45272fc0737e3420c0aa2c5c30a74f7a36",
|
||||
"commit": "8727258f04be6647d7b9013993d1d335acf84169",
|
||||
"url": "https://github.com/cubiq/ComfyUI_essentials.git"
|
||||
},
|
||||
{
|
||||
|
||||
@@ -1,2 +1,6 @@
|
||||
pydantic~=2.6.4
|
||||
jsonpath-ng~=1.6.1
|
||||
jsonpath-ng~=1.6.1
|
||||
imageio==2.34.2
|
||||
imageio-ffmpeg==0.5.1
|
||||
websockets==12.0
|
||||
httpx==0.27.0
|
||||
@@ -360,6 +360,10 @@ def get_history(prompt_id):
|
||||
with urllib.request.urlopen("http://{}/history/{}".format(SERVER_IP, prompt_id)) as response:
|
||||
return json.loads(response.read())
|
||||
|
||||
def get_queue(server_ip):
|
||||
with urllib.request.urlopen("http://{}/queue".format(server_ip)) as response:
|
||||
return json.loads(response.read())
|
||||
|
||||
async def get_images(ws, prompt_id, filter_node_id=None):
|
||||
output_images = {}
|
||||
while True:
|
||||
|
||||
Reference in New Issue
Block a user