This commit is contained in:
salt
2024-03-27 01:46:13 -07:00
parent 6c83155afc
commit ff52e26597
3 changed files with 165 additions and 10 deletions
Binary file not shown.
+148 -10
View File
@@ -1,18 +1,20 @@
import os
import signal
import subprocess
import time
from fastapi import FastAPI, File, UploadFile, HTTPException
from fastapi.responses import FileResponse
import shutil
from pathlib import Path
import asyncio
import uuid
import os
import json
import requests
import random
from typing import Dict, Any, Optional
from fastapi import FastAPI, File, UploadFile, HTTPException, Form
from fastapi.responses import FileResponse
import requests
from jsonpath_ng import jsonpath, parse
app = FastAPI()
lock = asyncio.Lock() # Create an asyncio lock
@@ -20,11 +22,53 @@ lock = asyncio.Lock() # Create an asyncio lock
SERVER_IP = os.environ.get("SERVER_IP", "127.0.0.1:8188")
TIMEOUT_SECONDS = os.environ.get("TIMEOUT_SECONDS", 1000)
PROMPT_ENDPOINT = f"http://{SERVER_IP}/prompt"
PROMPT_FILE = os.environ.get("PROMPT_FILE", "prompt.json")
V2V_WORKFLOWS_DIRECTORY = os.environ.get("V2V_WORKFLOWS_DIRECTORY",
"/home/salt/clone/storyteller/V2V-StylesList/V2VWorkflows/")
COMFY_ROOT_DIRECTORY = os.environ.get("COMFY_ROOT_DIRECTORY",
"/home/salt/clone/storyteller/storyteller-ml/workflows/comfy/ComfyUI")
SERVER_START_TIMEOUT = int(os.environ.get("SERVER_START_TIMEOUT", 60))
STYLE_LOCK = os.environ.get("STYLE_LOCK", None)
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",
}
allowed_keys = ["style", "positive_prompt", "negative_prompt"]
def validate_request(request):
try:
request = json.loads(request)
except json.JSONDecodeError:
raise HTTPException(status_code=400, detail="Invalid JSON for 'request'")
# check for style key
if "style" not in request:
raise HTTPException(status_code=400, detail="Missing 'style' key in 'request'")
style = request["style"]
if style not in style_to_filename:
raise HTTPException(status_code=400,
detail=f"Invalid 'style' key in 'request'. Valid styles: {', '.join(style_to_filename.keys())}")
# check for other keys
for key in request:
if key not in allowed_keys:
raise HTTPException(status_code=400,
detail=f"Invalid key '{key}' in 'request'. Valid keys: {', '.join(allowed_keys)}")
return request
def terminate_process(proc):
@@ -42,9 +86,16 @@ def signal_handler(signum, frame):
@app.post("/preview/")
async def upload_video(video: UploadFile = File(...)):
async def preview_request(request: str = Form(...), video: UploadFile = File(...)):
request = validate_request(request)
if STYLE_LOCK:
if request["style"] != STYLE_LOCK:
raise HTTPException(status_code=400, detail=f"This endpoint is for '{STYLE_LOCK}' style only")
async with lock: # Acquire the lock
try:
style_filename = style_to_filename[request["style"]]
positive_prompt = request.get("positive_prompt", None)
negative_prompt = request.get("negative_prompt", None)
# Save uploaded video to a temporary directory
video_filename = f"temp_videos/{uuid.uuid4()}.mp4" # Use a unique filename
Path("temp_videos").mkdir(parents=True, exist_ok=True) # Ensure the directory exists
@@ -54,16 +105,32 @@ async def upload_video(video: UploadFile = File(...)):
raise HTTPException(status_code=500, detail=f"Failed to save video: {str(e)}")
# Process video and generate preview
preview_image_path = await process_video(video_filename)
preview_image_path = await process_video(video_filename, style_filename, positive_prompt, negative_prompt)
return FileResponse(preview_image_path, media_type='image/jpeg')
async def process_video(video_path: str) -> Path:
json_prompt = json.load(open(PROMPT_FILE))
async def process_video(video_path: str, style_filename, positive_prompt=None, negative_prompt=None) -> Path:
# generate random seed (otherwise video will cache completely and not generate a new output)
styles_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "styles"
workflow_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "workflows"
mappings_directory = Path(V2V_WORKFLOWS_DIRECTORY) / "mappings"
style_json = json.load(open(styles_directory / style_filename))
workflow_json = json.load(open(workflow_directory / style_json["workflow_api_name"]))
mapping_json = json.load(open(mappings_directory / style_json["mapping_name"]))
json_mods = style_json["modifications"]
jsonpath_mods = get_jsonpath_mods(json_mods, mapping_json, pos_in=positive_prompt, neg_in=negative_prompt)
print(*jsonpath_mods.items(), sep="\n")
json_prompt = apply_jsonpath_mods(jsonpath_mods, workflow_json)
# generate random seed (otherwise video will cache completely and not generate a new output)
json_prompt["173"]["inputs"]["seed"] = random.randint(0, 1000000)
try:
json_prompt["173"]["inputs"]["seed"] = random.randint(0, 1000000)
except KeyError:
pass
input_video_path = Path(COMFY_ROOT_DIRECTORY) / "input/input.mp4"
preview_image_path = Path(COMFY_ROOT_DIRECTORY) / "output/vid2vid/SparseImagePreview_00001_.png"
@@ -116,8 +183,79 @@ async def process_video(video_path: str) -> Path:
return preview_image_path
def get_jsonpath_mods(style_mods: Dict[str, Any], mapping_json: Dict[str, Any], pos_in: Optional[str] = None,
neg_in: Optional[str] = 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:
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"{new_mod_json.get('positive_prompt', '')}, {pos_in}"
if neg_in:
new_mod_json["negative_prompt"] = f"{new_mod_json.get('negative_prompt', '')}, {neg_in}"
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}'")
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
if __name__ == "__main__":
import uvicorn
import asyncio
# install ComfyUI
os.system("python3 install.py")
signal.signal(signal.SIGINT, signal_handler)
signal.signal(signal.SIGTERM, signal_handler)
+17
View File
@@ -0,0 +1,17 @@
import json
import requests
url = 'http://localhost:8000/preview/'
payload = {"style": "comic_book", "positive_prompt": "test pos in", "negative_prompt": "test neg in"}
files = {'video': open('miku1f.mp4', 'rb')}
response = requests.post(url, data={'request': json.dumps(payload)}, files=files) # Use 'request': payload directly
if response.status_code == 200:
# Success
with open('preview_image.jpeg', 'wb') as f:
f.write(response.content)
else:
# Handle request error
print(f"Request failed: {response.status_code}")
print(response.text)