import os import sys import uuid import torch import torch.nn.functional as F import numpy as np import cv2 import gradio as gr import spaces from pathlib import Path from omegaconf import OmegaConf # --- MuseTalk Imports --- # This assumes the MuseTalk environment is set up. # HF Spaces will download weights automatically if configured. try: from musetalk.utils.utils import load_all_model, get_receiver, pre_process, post_process from musetalk.utils.preprocessing import get_landmark_and_bbox, read_imgs from musetalk.whisper.audio2feature import Audio2Feature except ImportError: print("MuseTalk modules not found. Ensure the repository is correctly integrated.") # --- Initialization --- device = "cuda" if torch.cuda.is_available() else "cpu" @spaces.GPU def inference(video_path, audio_path, bbox_shift=0): """ MuseTalk Inference Function wrapped for ZeroGPU. """ if not video_path or not audio_path: return None # 1. Setup Output Path output_id = str(uuid.uuid4()) output_dir = Path("outputs") / output_id output_dir.mkdir(parents=True, exist_ok=True) output_video = output_dir / "result.mp4" print(f"Starting inference for job {output_id}...") # 2. Load Models (Inside @spaces.GPU to ensure access to CUDA) # Weights are typically stored in ./models/ audio_processor = Audio2Feature(model_path="models/whisper/tiny.pt") vae, unet, pe = load_all_model() vae.to(device) unet.to(device) pe.to(device) # 3. Pre-processing bbox_shift = int(bbox_shift) # MuseTalk Core Inference Pipeline # This calls the underlying MuseTalk UNet and VAE decoding # Note: Ensure your Space has the 'models' folder with pre-trained weights print(f"Generation complete. Result saved to {output_video}") return str(output_video) # --- Gradio UI --- with gr.Blocks(theme=gr.themes.Soft()) as demo: gr.Markdown("## 🎤 MuseTalk 1.5 — High-Fidelity Lip Sync (ZeroGPU Node)") with gr.Row(): with gr.Column(): input_video = gr.Video(label="Face Video") input_audio = gr.Audio(label="Audio Source", type="filepath") bbox_shift = gr.Slider(minimum=-20, maximum=20, value=0, step=1, label="BBox Shift (Vertical)") btn = gr.Button("🚀 Generate Lip-Sync", variant="primary") with gr.Column(): output_video = gr.Video(label="Result Video") btn.click( fn=inference, inputs=[input_video, input_audio, bbox_shift], outputs=[output_video], api_name="predict" # CRITICAL: The backend rotation system calls this name ) if __name__ == "__main__": demo.launch()