dosesnrolls1's picture
Update app.py
e04c798 verified
Raw
History Blame Contribute Delete
4.71 kB
import os
import tempfile
import gradio as gr
import numpy as np
from PIL import Image
from composer import compose_frames, crop_reserved_region
from pipeline import load_pipeline, run_inference
from video_utils import (
compute_target_size,
extract_audio,
frames_for_duration,
load_video_frames,
resize_frames,
save_video,
)
DEFAULT_RESOLUTION = 768
REGION_SIZE = 256
_pipeline_state = None
def make_temp_file(suffix: str) -> str:
fd, path = tempfile.mkstemp(suffix=suffix)
os.close(fd)
return path
def generate(
guide_video_path,
face_image,
prompt,
duration,
fps,
lora_strength,
seed,
hf_token="",
progress=gr.Progress(track_tqdm=True),
):
global _pipeline_state
if not guide_video_path:
return None, "Please upload a guide video."
if face_image is None:
return None, "Please upload a reference face image."
if not prompt or not str(prompt).strip():
return None, "Please enter a text prompt."
if not os.path.isfile(guide_video_path):
return None, f"Guide video path is not a real file: {guide_video_path}"
if _pipeline_state is None:
progress(0, desc="Loading models (first run only — ~5 min)…")
_pipeline_state = load_pipeline(
token=str(hf_token).strip() or None,
progress_cb=lambda msg: progress(0, desc=msg),
)
progress(0.05, desc="Loading guide video…")
frames, source_fps = load_video_frames(guide_video_path)
if len(frames) == 0:
return None, "Could not read frames from the guide video."
audio_tmp = make_temp_file(".wav")
has_audio = extract_audio(guide_video_path, audio_tmp)
progress(0.10, desc="Resizing frames…")
orig_h, orig_w = frames.shape[1], frames.shape[2]
target_w, target_h = compute_target_size(orig_w, orig_h, DEFAULT_RESOLUTION)
frames = resize_frames(frames, target_w, target_h)
n_frames = frames_for_duration(fps, duration)
if len(frames) >= n_frames:
frames = frames[:n_frames]
else:
pad = np.stack([frames[-1]] * (n_frames - len(frames)))
frames = np.concatenate([frames, pad], axis=0)
progress(0.15, desc="Compositing reference face strip…")
composed = compose_frames(
frames,
face_image,
region_position="left",
region_size_px=REGION_SIZE,
)
progress(0.20, desc="Running LTX-2.3 diffusion…")
generated = run_inference(
_pipeline_state,
composed,
prompt=prompt,
fps=fps,
lora_strength=lora_strength,
seed=int(seed),
progress_cb=lambda msg: progress(0.20, desc=msg),
)
progress(0.90, desc="Cropping reserved region…")
cropped = crop_reserved_region(
generated,
region_position="left",
region_size_px=REGION_SIZE,
output_size=(target_w, target_h),
)
progress(0.95, desc="Encoding output video…")
out_path = make_temp_file(".mp4")
save_video(
cropped,
fps=fps,
output_path=out_path,
audio_path=audio_tmp if has_audio else None,
audio_duration=duration,
)
if not os.path.isfile(out_path):
return None, f"Output file was not created: {out_path}"
progress(1.0, desc="Done.")
return out_path, "Generation complete."
with gr.Blocks() as demo:
gr.Markdown("# BFS Best Face Swap Video")
with gr.Row():
guide_video = gr.Video(
label="Guide Video",
sources=["upload"],
)
face_image = gr.Image(
label="Reference Face Image",
type="pil",
)
prompt = gr.Textbox(
label="Prompt",
lines=4,
placeholder="headswap FACE ... ACTION ...",
)
with gr.Row():
duration = gr.Slider(1, 10, value=5, step=0.5, label="Duration (seconds)")
fps = gr.Slider(8, 30, value=24, step=1, label="FPS")
with gr.Row():
lora_strength = gr.Slider(0.0, 2.0, value=1.0, step=0.05, label="LoRA Strength")
seed = gr.Number(value=42, precision=0, label="Seed")
hf_token = gr.Textbox(
label="Hugging Face Token (optional)",
type="password",
)
run_btn = gr.Button("Generate", variant="primary")
output_video = gr.Video(label="Output Video")
status = gr.Textbox(label="Status", interactive=False)
run_btn.click(
fn=generate,
inputs=[
guide_video,
face_image,
prompt,
duration,
fps,
lora_strength,
seed,
hf_token,
],
outputs=[output_video, status],
)
demo.launch()