Spaces:
Running on Zero
Running on Zero
| #!/usr/bin/env python3 | |
| """Gradio demo for InstructAV2AV: source video + instruction -> edited video.""" | |
| from __future__ import annotations | |
| import argparse | |
| import gc | |
| import logging | |
| import os | |
| import sys | |
| import uuid | |
| from collections import OrderedDict | |
| from pathlib import Path | |
| from threading import Lock | |
| from typing import Any, Callable | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| if str(REPO_ROOT) not in sys.path: | |
| sys.path.insert(0, str(REPO_ROOT)) | |
| # Some upstream Ovi modules resolve their default config relative to cwd. | |
| os.chdir(REPO_ROOT) | |
| import gradio as gr | |
| import torch | |
| from omegaconf import OmegaConf | |
| from ovi.distributed_comms.parallel_states import initialize_sequence_parallel_state | |
| from ovi.utils.av_edit_data import ( | |
| get_video_info, | |
| load_audio_array, | |
| load_video_array, | |
| snap_num_frames, | |
| to_audio_tensor, | |
| to_video_tensor, | |
| ) | |
| from ovi.utils.io_utils import save_video | |
| DEFAULT_CONFIG = REPO_ROOT / "ovi/configs/inference/inference_av_edit.yaml" | |
| DEFAULT_MODEL_DIR = REPO_ROOT / "ckpts/InstructAV2AV" | |
| DEFAULT_OUTPUT_DIR = REPO_ROOT / "outputs/demo" | |
| MODEL_SPECS = OrderedDict( | |
| { | |
| "general": { | |
| "label": "General Edit", | |
| "description": "Flexible editing of appearance, scenes, actions, speech, and sound.", | |
| "example": "Make the horse dark brown with a white saddle.", | |
| }, | |
| "insertion": { | |
| "label": "Content Insertion", | |
| "description": "Add a an object to the source video.", | |
| "example": "Add a dark vintage sedan driving from the right to the left.", | |
| }, | |
| "removal": { | |
| "label": "Content Removal", | |
| "description": "Remove an object from the video.", | |
| "example": "Remove the chipmunk standing on the stone surface among the peanuts.", | |
| }, | |
| "clone_id": { | |
| "label": "Identity Cloning", | |
| "description": "Preserve a person's visual identity during editing.", | |
| "example": "Keep the person‘s appearance, change the timbre to a man, and change the spoken words to <S>I understand, but I think we need to consider.<E>.", | |
| }, | |
| "clone_voice": { | |
| "label": "Voice Cloning", | |
| "description": "Preserve the speaker's timbre during editing.", | |
| "example": "Keep the timbre, change ..., and change the spoken words to <S>I came here to tell you that you should to go.<E>.", | |
| }, | |
| "clone_id_voice": { | |
| "label": "Identity + Voice Cloning", | |
| "description": "Preserve both visual identity and timbre.", | |
| "example": "Keep the person’s identity and change the spoken words to <S>This is more than just art, it’s a statement.<E>.", | |
| }, | |
| } | |
| ) | |
| CSS = """ | |
| .gradio-container { | |
| max-width: 1200px !important; | |
| margin: 0 auto !important; | |
| padding: 24px 20px 32px !important; | |
| font-family: Arial, Helvetica, sans-serif !important; | |
| } | |
| #page-header { margin-bottom: 18px !important; } | |
| #page-header h1 { margin-bottom: 4px !important; font-size: 1.7rem !important; } | |
| #page-header p { margin: 0 !important; color: var(--body-text-color-subdued); } | |
| #video-row { gap: 18px !important; align-items: start !important; } | |
| #source-column, #result-column { gap: 8px !important; min-width: 0 !important; } | |
| #source-video, #result-video { | |
| height: auto !important; | |
| margin: 0 !important; | |
| aspect-ratio: 16 / 9; | |
| } | |
| #source-video [data-testid="video"], | |
| #result-video [data-testid="video"] { | |
| width: 100% !important; | |
| height: auto !important; | |
| aspect-ratio: 16 / 9; | |
| overflow: hidden; | |
| } | |
| #source-video video, #result-video video { | |
| width: 100% !important; | |
| height: 100% !important; | |
| object-fit: contain !important; | |
| background: #000 !important; | |
| border-radius: 4px !important; | |
| } | |
| #settings-panel { | |
| margin-top: 18px !important; | |
| padding: 16px !important; | |
| border: 1px solid var(--border-color-primary) !important; | |
| border-radius: 6px !important; | |
| box-shadow: none !important; | |
| } | |
| #settings-title { margin: 0 0 6px !important; } | |
| #settings-title h2 { margin: 0 !important; font-size: 1.1rem !important; } | |
| #settings-row { align-items: start !important; gap: 14px !important; } | |
| #model-help { | |
| margin: 0 !important; | |
| padding: 2px !important; | |
| color: var(--body-text-color-subdued); | |
| font-size: 0.83rem; | |
| } | |
| #model-help .prose { padding: 0 !important; } | |
| #model-help p { margin: 0 0 4px !important; } | |
| #model-help p:last-child { margin-bottom: 0 !important; } | |
| #advanced-settings { margin-top: 6px !important; } | |
| #action-row { gap: 10px !important; justify-content: flex-end !important; } | |
| #generate-button { | |
| min-height: 40px; | |
| border-radius: 4px; | |
| } | |
| #clear-button { min-height: 40px; border-radius: 4px; } | |
| @media (max-width: 760px) { | |
| .gradio-container { padding: 10px !important; } | |
| #video-row { flex-direction: column !important; } | |
| #source-column, #result-column { width: 100% !important; } | |
| #video-row, #settings-row { gap: 10px !important; } | |
| #settings-panel { padding: 12px !important; } | |
| } | |
| """ | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config-file", default=str(DEFAULT_CONFIG)) | |
| parser.add_argument("--model-dir", default=str(DEFAULT_MODEL_DIR)) | |
| parser.add_argument("--ckpt-dir", default=None, help="Base Ovi checkpoint directory override.") | |
| parser.add_argument("--output-dir", default=str(DEFAULT_OUTPUT_DIR)) | |
| parser.add_argument("--device", type=int, default=0) | |
| parser.add_argument("--server-name", default="127.0.0.1") | |
| parser.add_argument("--server-port", type=int, default=7860) | |
| parser.add_argument("--share", action="store_true") | |
| parser.add_argument("--inbrowser", action="store_true") | |
| parser.add_argument("--no-cpu-offload", action="store_true") | |
| parser.add_argument("--max-queue-size", type=int, default=8) | |
| return parser.parse_args() | |
| def resolve_path(value: str | Path, base: Path = REPO_ROOT) -> Path: | |
| path = Path(value).expanduser() | |
| if not path.is_absolute(): | |
| path = base / path | |
| return path.resolve() | |
| def discover_checkpoints(model_dir: Path) -> dict[str, Path]: | |
| checkpoints = {key: model_dir / f"{key}.safetensors" for key in MODEL_SPECS} | |
| missing = [str(path) for path in checkpoints.values() if not path.is_file()] | |
| if missing: | |
| raise FileNotFoundError( | |
| "Missing InstructAV2AV checkpoints:\n" | |
| + "\n".join(f"- {path}" for path in missing) | |
| ) | |
| return checkpoints | |
| def uploaded_path(value: Any) -> Path: | |
| if value is None: | |
| raise gr.Error("Please upload a source video first.") | |
| if isinstance(value, (str, Path)): | |
| path = Path(value) | |
| elif isinstance(value, dict): | |
| raw_path = value.get("path") or value.get("name") | |
| if not raw_path: | |
| raise gr.Error("The uploaded video could not be read.") | |
| path = Path(raw_path) | |
| else: | |
| raise gr.Error(f"Unsupported video input type: {type(value).__name__}") | |
| path = path.expanduser().resolve() | |
| if not path.is_file(): | |
| raise gr.Error(f"The uploaded video does not exist: {path}") | |
| return path | |
| class DemoRuntime: | |
| """Keep shared modules alive and swap only the selected fusion checkpoint.""" | |
| def __init__( | |
| self, | |
| args: argparse.Namespace, | |
| checkpoint_resolver: Callable[[str, Any], Path] | None = None, | |
| ): | |
| self.args = args | |
| self.config_path = resolve_path(args.config_file) | |
| self.model_dir = resolve_path(args.model_dir) | |
| self.output_dir = resolve_path(args.output_dir) | |
| self.checkpoint_resolver = checkpoint_resolver | |
| self.checkpoints = ( | |
| {} | |
| if checkpoint_resolver is not None | |
| else discover_checkpoints(self.model_dir) | |
| ) | |
| self.config = self._load_config() | |
| self.engine: Any | None = None | |
| self.active_model: str | None = None | |
| self.lock = Lock() | |
| self.output_dir.mkdir(parents=True, exist_ok=True) | |
| def _load_config(self): | |
| if not self.config_path.is_file(): | |
| raise FileNotFoundError(f"Inference config not found: {self.config_path}") | |
| config = OmegaConf.load(self.config_path) | |
| ckpt_dir = self.args.ckpt_dir or config.get("ckpt_dir", "./ckpts") | |
| config.ckpt_dir = str(resolve_path(ckpt_dir)) | |
| config.av2av_edit = True | |
| config.has_video = True | |
| config.has_audio = True | |
| config.mode = "t2v" | |
| config.sp_size = 1 | |
| config.cpu_offload = not self.args.no_cpu_offload | |
| return config | |
| def _validate_cuda(self) -> None: | |
| if not torch.cuda.is_available(): | |
| raise RuntimeError("InstructAV2AV inference requires a CUDA GPU.") | |
| device_count = torch.cuda.device_count() | |
| if self.args.device < 0 or self.args.device >= device_count: | |
| raise RuntimeError( | |
| f"CUDA device {self.args.device} is unavailable; " | |
| f"found {device_count} GPU(s)." | |
| ) | |
| def get_engine(self, model_key: str, progress: gr.Progress) -> Any: | |
| if model_key not in MODEL_SPECS: | |
| raise ValueError(f"Unknown model: {model_key}") | |
| self._validate_cuda() | |
| # Keep the web page startup light; import the large model stack on first use. | |
| from ovi.ovi_fusion_engine import OviFusionEngine | |
| from ovi.utils.model_loading_utils import load_fusion_checkpoint | |
| if self.checkpoint_resolver is None: | |
| checkpoint = self.checkpoints[model_key] | |
| else: | |
| checkpoint = Path(self.checkpoint_resolver(model_key, progress)).resolve() | |
| if not checkpoint.is_file(): | |
| raise FileNotFoundError(f"Editing checkpoint not found: {checkpoint}") | |
| if self.engine is None: | |
| progress(0.08, desc=f"Loading {MODEL_SPECS[model_key]['label']} model") | |
| torch.cuda.set_device(self.args.device) | |
| initialize_sequence_parallel_state(1) | |
| config = OmegaConf.create(OmegaConf.to_container(self.config, resolve=True)) | |
| config.finetune_path = str(checkpoint) | |
| self.engine = OviFusionEngine( | |
| config=config, | |
| device=self.args.device, | |
| target_dtype=torch.bfloat16, | |
| ).eval() | |
| self.active_model = model_key | |
| elif self.active_model != model_key: | |
| progress(0.08, desc=f"Switching to {MODEL_SPECS[model_key]['label']}…") | |
| # All six checkpoints share one architecture, so the T5 and VAEs stay loaded. | |
| self.active_model = None | |
| self.engine.model = self.engine.model.to("cpu") | |
| torch.cuda.empty_cache() | |
| load_fusion_checkpoint(self.engine.model, str(checkpoint), from_meta=False) | |
| if not self.engine.cpu_offload: | |
| self.engine.model = self.engine.model.to(device=self.args.device) | |
| self.engine.eval() | |
| self.active_model = model_key | |
| gc.collect() | |
| return self.engine | |
| def generate( | |
| self, | |
| video_value: Any, | |
| instruction: str, | |
| model_key: str, | |
| seed: float, | |
| sample_steps: float, | |
| video_guidance_scale: float, | |
| audio_guidance_scale: float, | |
| progress: gr.Progress = gr.Progress(), | |
| ) -> str: | |
| video_path = uploaded_path(video_value) | |
| instruction = (instruction or "").strip() | |
| if not instruction: | |
| raise gr.Error("Please enter an editing instruction.") | |
| if model_key not in MODEL_SPECS: | |
| raise gr.Error("Please select a valid model.") | |
| sample_steps = int(sample_steps) | |
| seed = int(seed) | |
| if not 1 <= sample_steps <= 100: | |
| raise gr.Error("Sampling steps must be between 1 and 100.") | |
| try: | |
| with self.lock: | |
| engine = self.get_engine(model_key, progress) | |
| progress(0.18, desc="Preparing the source video and audio…") | |
| fps, total_frames = get_video_info(video_path) | |
| configured_frames = int(self.config.get("num_frames", total_frames)) | |
| num_frames = snap_num_frames(min(total_frames, configured_frames)) | |
| frame_size = list(self.config.get("video_frame_height_width", [704, 1280])) | |
| video, _ = load_video_array( | |
| video_path, | |
| num_frames=num_frames, | |
| height=int(frame_size[0]), | |
| width=int(frame_size[1]), | |
| max_pixels=int(frame_size[0]) * int(frame_size[1]), | |
| ) | |
| sample_rate = int(self.config.get("audio_sample_rate", 16000)) | |
| audio_samples = max(1, round(num_frames / fps * sample_rate)) | |
| try: | |
| audio = load_audio_array( | |
| video_path, | |
| sample_rate=sample_rate, | |
| num_samples=audio_samples, | |
| ) | |
| except Exception as exc: | |
| raise RuntimeError( | |
| "The video's audio track could not be read. " | |
| "Please upload a video that contains audio." | |
| ) from exc | |
| input_video = to_video_tensor(video, engine.device, engine.target_dtype) | |
| input_audio = to_audio_tensor(audio, engine.device) | |
| del video, audio | |
| progress(0.28, desc="Generating the edit. This may take several minutes…") | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed_all(seed) | |
| generated = engine.generate( | |
| text_prompt=instruction, | |
| image_path=None, | |
| video_frame_height_width=frame_size, | |
| seed=seed, | |
| solver_name=str(self.config.get("solver_name", "unipc")), | |
| sample_steps=sample_steps, | |
| shift=float(self.config.get("shift", 5.0)), | |
| video_guidance_scale=float(video_guidance_scale), | |
| audio_guidance_scale=float(audio_guidance_scale), | |
| slg_layer=int(self.config.get("slg_layer", 11)), | |
| video_negative_prompt=str(self.config.get("video_negative_prompt", "")), | |
| audio_negative_prompt=str(self.config.get("audio_negative_prompt", "")), | |
| input_video=input_video, | |
| input_audio=input_audio, | |
| ) | |
| del input_video, input_audio | |
| if generated is None: | |
| raise RuntimeError( | |
| "Generation failed. Check the terminal log for details." | |
| ) | |
| generated_video, generated_audio, _ = generated | |
| if generated_video is None or generated_audio is None: | |
| raise RuntimeError("The model returned an incomplete audio-video result.") | |
| progress(0.92, desc="Encoding the edited video…") | |
| output_path = self.output_dir / f"{model_key}_{uuid.uuid4().hex[:12]}.mp4" | |
| save_video( | |
| str(output_path), | |
| generated_video, | |
| generated_audio, | |
| sample_rate=sample_rate, | |
| fps=fps, | |
| ) | |
| del generated_video, generated_audio, generated | |
| gc.collect() | |
| torch.cuda.empty_cache() | |
| progress(1.0, desc="Edit complete") | |
| return str(output_path) | |
| except gr.Error: | |
| raise | |
| except Exception as exc: | |
| logging.exception("InstructAV2AV demo generation failed") | |
| raise gr.Error(str(exc)) from exc | |
| def warmup( | |
| self, | |
| model_key: str = "general", | |
| progress: gr.Progress = gr.Progress(), | |
| ) -> str: | |
| """Initialize the default engine before the first edit request.""" | |
| try: | |
| with self.lock: | |
| self.get_engine(model_key, progress) | |
| return f"✅ {MODEL_SPECS[model_key]['label']} is ready on ZeroGPU." | |
| except Exception as exc: | |
| logging.exception("InstructAV2AV demo warmup failed") | |
| return f"⚠️ Model warmup did not finish: {exc}" | |
| def model_help(model_key: str) -> str: | |
| spec = MODEL_SPECS.get(model_key, MODEL_SPECS["general"]) | |
| return f"{spec['description']} \n**Example:** `{spec['example']}`" | |
| def build_demo( | |
| runtime: DemoRuntime, | |
| generate_fn: Callable[..., str] | None = None, | |
| warmup_fn: Callable[..., str] | None = None, | |
| zero_gpu_fn: Callable[..., str] | None = None, | |
| ) -> gr.Blocks: | |
| choices = [(spec["label"], key) for key, spec in MODEL_SPECS.items()] | |
| config = runtime.config | |
| with gr.Blocks(title="InstructAV2AV Demo") as demo: | |
| warmup_operation = gr.State("warmup") | |
| generate_operation = gr.State("generate") | |
| gr.Markdown( | |
| "# InstructAV2AV\nUpload a video with audio, choose an editing type, provide the edit instruction, and generate the result. The default General model is prepared when the Space starts.", | |
| elem_id="page-header", | |
| ) | |
| model_status = gr.Markdown( | |
| "⏳ Preparing the default General model…", | |
| elem_id="model-status", | |
| ) | |
| with gr.Row(equal_height=True, elem_id="video-row"): | |
| with gr.Column(scale=1, min_width=0, elem_id="source-column"): | |
| source_video = gr.Video( | |
| label="Source video", | |
| sources=["upload"], | |
| include_audio=True, | |
| elem_id="source-video", | |
| elem_classes=["video-card"], | |
| ) | |
| with gr.Column(scale=1, min_width=0, elem_id="result-column"): | |
| result_video = gr.Video( | |
| label="Edited video", | |
| format="mp4", | |
| interactive=False, | |
| elem_id="result-video", | |
| elem_classes=["video-card"], | |
| ) | |
| with gr.Group(elem_id="settings-panel"): | |
| gr.Markdown("## Edit settings", elem_id="settings-title") | |
| with gr.Row(elem_id="settings-row"): | |
| with gr.Column(scale=1, min_width=260): | |
| model_choice = gr.Dropdown( | |
| choices=choices, | |
| value="general", | |
| label="Editing type", | |
| allow_custom_value=False, | |
| ) | |
| model_description = gr.Markdown( | |
| model_help("general"), elem_id="model-help" | |
| ) | |
| instruction = gr.Textbox( | |
| label="Editing instruction", | |
| placeholder="Example: Change the man into a young woman with brown hair, wearing a gray blazer, and saying, <S>I really think we should give it another chance.<E>.", | |
| info=( | |
| "For speech editing, wrap the spoken text with <S> and <E>." | |
| ), | |
| lines=4, | |
| max_lines=8, | |
| scale=2, | |
| min_width=320, | |
| ) | |
| with gr.Accordion("Advanced settings", open=False, elem_id="advanced-settings"): | |
| with gr.Row(): | |
| seed = gr.Number( | |
| value=int(config.get("seed", 103)), | |
| label="Seed", | |
| precision=0, | |
| ) | |
| sample_steps = gr.Slider( | |
| minimum=1, | |
| maximum=100, | |
| value=int(config.get("sample_steps", 50)), | |
| step=1, | |
| label="Sampling steps", | |
| ) | |
| with gr.Row(): | |
| video_guidance = gr.Slider( | |
| minimum=0, | |
| maximum=10, | |
| value=float(config.get("video_guidance_scale", 4.0)), | |
| step=0.1, | |
| label="Video guidance", | |
| ) | |
| audio_guidance = gr.Slider( | |
| minimum=0, | |
| maximum=10, | |
| value=float(config.get("audio_guidance_scale", 3.0)), | |
| step=0.1, | |
| label="Audio guidance", | |
| ) | |
| with gr.Row(elem_id="action-row"): | |
| clear_button = gr.Button( | |
| "Clear", variant="secondary", elem_id="clear-button" | |
| ) | |
| generate_button = gr.Button( | |
| "Generate edit", | |
| variant="primary", | |
| elem_id="generate-button", | |
| ) | |
| model_choice.change( | |
| model_help, | |
| inputs=model_choice, | |
| outputs=model_description, | |
| queue=False, | |
| ) | |
| generation_inputs = [ | |
| source_video, | |
| instruction, | |
| model_choice, | |
| seed, | |
| sample_steps, | |
| video_guidance, | |
| audio_guidance, | |
| ] | |
| generate_button.click( | |
| zero_gpu_fn or generate_fn or runtime.generate, | |
| inputs=( | |
| [generate_operation, *generation_inputs] | |
| if zero_gpu_fn is not None | |
| else generation_inputs | |
| ), | |
| outputs=result_video, | |
| concurrency_limit=1, | |
| concurrency_id="instructav2av-generation", | |
| api_name="edit_video", | |
| ) | |
| clear_button.click( | |
| lambda: ( | |
| None, | |
| "", | |
| None, | |
| ), | |
| outputs=[source_video, instruction, result_video], | |
| queue=False, | |
| ) | |
| if zero_gpu_fn is not None: | |
| demo.load( | |
| zero_gpu_fn, | |
| inputs=[warmup_operation, *generation_inputs], | |
| outputs=model_status, | |
| concurrency_limit=1, | |
| concurrency_id="instructav2av-generation", | |
| ) | |
| elif warmup_fn is not None: | |
| demo.load( | |
| warmup_fn, | |
| outputs=model_status, | |
| concurrency_limit=1, | |
| concurrency_id="instructav2av-generation", | |
| ) | |
| return demo | |
| def main() -> None: | |
| args = parse_args() | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="[%(asctime)s] %(levelname)s: %(message)s", | |
| ) | |
| runtime = DemoRuntime(args) | |
| demo = build_demo(runtime) | |
| demo.queue(max_size=args.max_queue_size, default_concurrency_limit=1) | |
| demo.launch( | |
| server_name=args.server_name, | |
| server_port=args.server_port, | |
| share=args.share, | |
| inbrowser=args.inbrowser, | |
| allowed_paths=[str(runtime.output_dir)], | |
| show_error=True, | |
| theme=gr.themes.Default(), | |
| css=CSS, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |