r3gm's picture
Update app.py
65be70c verified
Raw
History Blame Contribute Delete
56.3 kB
import os
os.system("pip list")
import spaces
import transformers.utils.import_utils as transformers_import_utils
# https://huggingface.co/spaces/zero-gpu-explorers/README/discussions/181#6a711fa1aed3dcd9708e4b06
if hasattr(transformers_import_utils, "is_cuda_stream_capturing"):
_original_is_cuda_stream_capturing = (
transformers_import_utils.is_cuda_stream_capturing
)
def _zerogpu_safe_is_cuda_stream_capturing(*args, **kwargs):
if (
os.getenv("SPACES_ZERO_GPU") == "1"
and not os.getenv("CUDA_VISIBLE_DEVICES")
):
return False
return _original_is_cuda_stream_capturing(*args, **kwargs)
transformers_import_utils.is_cuda_stream_capturing = (
_zerogpu_safe_is_cuda_stream_capturing
)
import shutil
import subprocess
import sys
import copy
import random
import tempfile
import warnings
import time
import threading
import gc
import uuid
import re
import json
from tqdm import tqdm
import cv2
import numpy as np
import torch
import torch._dynamo
from huggingface_hub import list_models, hf_hub_download
from torch.nn import functional as F
from PIL import Image
import gradio as gr
from diffusers import (
FlowMatchEulerDiscreteScheduler,
SASolverScheduler,
DEISMultistepScheduler,
DPMSolverMultistepInverseScheduler,
UniPCMultistepScheduler,
DPMSolverMultistepScheduler,
DPMSolverSinglestepScheduler,
)
from diffusers.pipelines.wan.pipeline_wan_i2v import WanImageToVideoPipeline
from PIL.PngImagePlugin import PngInfo
from torchao.quantization import quantize_, Float8DynamicActivationFloat8WeightConfig, Int8WeightOnlyConfig
import aoti
from spandrel import ModelLoader
from safety import check_nsfw, pre_gpu_safety_check
os.environ["TOKENIZERS_PARALLELISM"] = "true"
warnings.filterwarnings("ignore")
IS_ZERO_GPU = bool(os.getenv("SPACES_ZERO_GPU"))
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# if IS_ZERO_GPU:
# print("Loading...")
# subprocess.run("rm -rf /data-nvme/zerogpu-offload/*", env={}, shell=True)
def cleanup_temp_dir(max_age_seconds=3600):
temp_dir = tempfile.gettempdir()
if not os.path.exists(temp_dir):
return
now = time.time()
valid_extensions = ('.mp4', '.png', '.json', '.txt')
try:
with os.scandir(temp_dir) as entries:
for entry in entries:
if entry.is_file():
if entry.name.lower().endswith(valid_extensions):
try:
if now - entry.stat().st_mtime >= max_age_seconds:
os.remove(entry.path)
except Exception:
pass
except Exception:
pass
def start_cleanup_daemon(interval_seconds=900, max_age_seconds=1200):
def daemon_loop():
while True:
time.sleep(interval_seconds)
cleanup_temp_dir(max_age_seconds)
cleanup_temp_dir(max_age_seconds)
t = threading.Thread(target=daemon_loop, daemon=True)
t.start()
start_cleanup_daemon()
css = """
button.primary-btn, button.generate-btn, .gradio-container button.primary {
width: 100% !important;
padding: 0.75rem !important;
background-color: var(--button-primary-background-fill, #111827) !important;
color: var(--button-primary-text-color, #ffffff) !important;
border: none !important;
border-radius: 4px !important;
font-size: 0.875rem !important;
font-weight: 500 !important;
cursor: pointer !important;
transition: background-color 0.15s ease !important;
display: inline-flex !important;
justify-content: center !important;
align-items: center !important;
gap: 0.5rem !important;
box-shadow: none !important;
}
button.primary-btn:hover, button.generate-btn:hover, .gradio-container button.primary:hover {
background-color: var(--button-primary-background-fill-hover, #374151) !important;
}
.dark button.primary-btn, .dark button.generate-btn, .dark .gradio-container button.primary {
background-color: var(--button-primary-background-fill, #f3f4f6) !important;
color: var(--button-primary-text-color, #111827) !important;
}
.dark button.primary-btn:hover, .dark button.generate-btn:hover, .dark .gradio-container button.primary:hover {
background-color: var(--button-primary-background-fill-hover, #e5e7eb) !important;
}
.compact-btn-row {
display: flex !important;
flex-direction: row !important;
align-items: center !important;
gap: 0.5rem !important;
flex-wrap: wrap !important;
margin-top: 0.5rem !important;
}
.compact-btn {
width: auto !important;
min-width: unset !important;
max-width: fit-content !important;
display: inline-flex !important;
padding: 4px 10px !important;
height: 32px !important;
font-size: 0.813rem !important;
}
:root {
--component-border-color: var(--neutral-200, rgba(128, 128, 128, 0.18));
--component-bg-color: var(--neutral-50, rgba(128, 128, 128, 0.03));
}
.dark {
--component-border-color: var(--neutral-800, rgba(255, 255, 255, 0.12));
--component-bg-color: var(--neutral-900, rgba(255, 255, 255, 0.03));
}
.gradio-container .block,
.gradio-container .panel,
.gradio-container .form,
.gradio-container fieldset {
border: 1px solid var(--component-border-color) !important;
background-color: var(--component-bg-color) !important;
border-radius: 8px !important;
box-shadow: none !important;
}
.gradio-container .markdown,
.gradio-container .prose,
.gradio-container .block.prose,
.gradio-container div[class*="markdown"] {
border: none !important;
background: transparent !important;
background-color: transparent !important;
box-shadow: none !important;
padding: 0 !important;
}
"""
# --- FRAME EXTRACTION JS & LOGIC ---
step_back_js = """
function() {
const video = document.querySelector('#generated-video video');
if (video) {
video.pause();
const fps = 16;
video.currentTime = Math.max(0, video.currentTime - (1.0 / fps));
}
return 0;
}
"""
step_fwd_js = """
function() {
const video = document.querySelector('#generated-video video');
if (video) {
video.pause();
const fps = 16;
video.currentTime = Math.min(video.duration || 0, video.currentTime + (1.0 / fps));
}
return 0;
}
"""
def extract_frame(video_path, timestamp):
# Safety check: if no video is present
if not video_path:
return None
if isinstance(video_path, dict):
video_path = video_path.get("video") or video_path.get("path")
elif hasattr(video_path, "name"):
video_path = video_path.name
print(f"Extracting frame at timestamp: {timestamp}")
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
return None
timestamp_ms = float(timestamp or 0.0) * 1000.0
cap.set(cv2.CAP_PROP_POS_MSEC, timestamp_ms)
ret, frame = cap.read()
if not ret:
fps = cap.get(cv2.CAP_PROP_FPS) or 16.0
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
target_frame_num = min(max(0, int(round(float(timestamp or 0.0) * fps))), total_frames - 1)
cap.set(cv2.CAP_PROP_POS_FRAMES, target_frame_num)
ret, frame = cap.read()
cap.release()
if ret:
# Convert from BGR (OpenCV) to RGB (PIL)
rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
pil_img = Image.fromarray(rgb_frame)
pnginfo = PngInfo()
pnginfo.add_text("parameters", json.dumps({"ai_generated": True, "source": "wan_extracted_frame"}))
pnginfo.add_text("ai_generated", "true")
pnginfo.add_text("prompt", "wan_video_generated_frame")
exif = pil_img.getexif()
exif[0x9286] = "ai_generated wan_extracted_frame"
exif[0x010e] = "ai_generated wan_extracted_frame"
# Save to temp file with persistent metadata so Gradio retains it
temp_file = tempfile.NamedTemporaryFile(suffix=".png", delete=False)
temp_file.close()
pil_img.save(temp_file.name, format="PNG", pnginfo=pnginfo, exif=exif)
return temp_file.name
return None
# --- END FRAME EXTRACTION LOGIC ---
def split_video_frame_level(video_path, timestamp):
if not video_path:
gr.Info("Please generate a video first.")
return None
if isinstance(video_path, dict):
video_path = video_path.get("video") or video_path.get("path")
elif hasattr(video_path, "name"):
video_path = video_path.name
import imageio_ffmpeg
ffmpeg_exe = imageio_ffmpeg.get_ffmpeg_exe()
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
gr.Warning("Unable to read video file.")
return None
fps = cap.get(cv2.CAP_PROP_FPS) or 16.0
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
cap.release()
split_frame = int(round(float(timestamp or 0.0) * fps))
if split_frame <= 0 or split_frame >= total_frames:
gr.Warning(f"Split point must be between start and end (Frame {split_frame}/{total_frames}). Pause the video during playback.")
return None
base_name = os.path.splitext(os.path.basename(video_path))[0]
out_dir = tempfile.gettempdir()
part1_path = os.path.join(out_dir, f"{base_name}_split_1.mp4")
part2_path = os.path.join(out_dir, f"{base_name}_split_2.mp4")
cmd1 = [
ffmpeg_exe, "-y", "-i", video_path,
"-vf", f"select='between(n,0,{split_frame - 1})',setpts=N/FRAME_RATE/TB",
"-c:v", "libx264", "-crf", "17", "-preset", "fast",
"-an", part1_path
]
cmd2 = [
ffmpeg_exe, "-y", "-i", video_path,
"-vf", f"select='between(n,{split_frame},{total_frames - 1})',setpts=N/FRAME_RATE/TB",
"-c:v", "libx264", "-crf", "17", "-preset", "fast",
"-an", part2_path
]
subprocess.run(cmd1, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
subprocess.run(cmd2, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
gr.Info(f"Video split at frame {split_frame}/{total_frames} ({float(timestamp):.2f}s).")
return [part1_path, part2_path]
def send_first_file_to_merger(file_data, current_list):
if not file_data:
return current_list or []
first_file = file_data[0] if isinstance(file_data, list) else file_data
first_path = first_file.name if hasattr(first_file, "name") else str(first_file)
existing = current_list or []
if not isinstance(existing, list):
existing = [existing]
paths = [f.name if hasattr(f, "name") else str(f) for f in existing]
if first_path not in paths:
paths.append(first_path)
return paths
def merge_videos(file_list):
if not file_list or len(file_list) < 2:
gr.Warning("Please add at least 2 videos to merge.")
return gr.update()
import imageio_ffmpeg
ffmpeg_exe = imageio_ffmpeg.get_ffmpeg_exe()
paths = [f.name if hasattr(f, "name") else str(f) for f in file_list]
out_dir = tempfile.gettempdir()
concat_txt = os.path.join(out_dir, f"concat_{uuid.uuid4().hex[:8]}.txt")
output_merged = os.path.join(out_dir, f"merged_{uuid.uuid4().hex[:8]}.mp4")
with open(concat_txt, "w", encoding="utf-8") as f:
for p in paths:
escaped_p = p.replace("'", "'\\''")
f.write(f"file '{escaped_p}'\n")
cmd = [
ffmpeg_exe, "-y", "-f", "concat", "-safe", "0",
"-i", concat_txt, "-c", "copy", output_merged
]
res = subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
if res.returncode != 0:
inputs = []
filter_str = ""
for i, p in enumerate(paths):
inputs.extend(["-i", p])
filter_str += f"[{i}:v:0]"
filter_str += f"concat=n={len(paths)}:v=1:a=0[v]"
cmd_fallback = [ffmpeg_exe, "-y", *inputs, "-filter_complex", filter_str, "-map", "[v]", "-c:v", "libx264", "-crf", "17", "-preset", "fast", output_merged]
subprocess.run(cmd_fallback, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
gr.Info(f"Successfully merged {len(paths)} videos!")
return output_merged
def export_to_video(
video_frames: list[np.ndarray] | list[Image.Image],
output_video_path: str = None,
fps: int = 10,
quality: float = 5.0,
bitrate: int | None = None,
macro_block_size: int | None = 16,
metadata: dict | None = None,
) -> str:
import imageio
if output_video_path is None:
output_video_path = tempfile.NamedTemporaryFile(suffix=".mp4").name
if isinstance(video_frames[0], np.ndarray):
video_frames = [(frame * 255).astype(np.uint8) for frame in video_frames]
elif isinstance(video_frames[0], Image.Image):
video_frames = [np.array(frame) for frame in video_frames]
output_params = []
if metadata is not None:
json_str = json.dumps(metadata, ensure_ascii=False)
output_params.extend(["-metadata", f"comment={json_str}"])
output_params.extend(["-metadata", "description=This content was generated by AI"])
output_params.extend(["-metadata", "artist=This content was generated by AI"])
num_frames_video = len(video_frames)
with (
imageio.get_writer(
output_video_path,
fps=fps,
quality=quality,
bitrate=bitrate,
macro_block_size=macro_block_size,
output_params=output_params if output_params else None,
) as writer,
tqdm(total=num_frames_video, desc="Encoding Video", unit="frame") as pbar
):
for i, frame in enumerate(video_frames):
writer.append_data(frame)
if (i + 1) % 25 == 0:
pbar.update(25)
pbar.update(num_frames_video % 25)
return output_video_path
def export_settings_json(
prompt, negative_prompt, steps, duration_seconds,
guidance_scale, guidance_scale_2, seed, randomize_seed, quality,
scheduler, flow_shift, frame_multiplier, safe_mode, safety_checker,
play_result_video, upscale_model, upscale_factor
):
"""
Serializes current UI settings into a downloadable JSON file path for gr.DownloadButton.
"""
settings_dict = {
"prompt": prompt,
"negative_prompt": negative_prompt,
"steps": int(steps),
"duration_seconds": float(duration_seconds),
"guidance_scale": float(guidance_scale),
"guidance_scale_2": float(guidance_scale_2),
"seed": int(seed),
"randomize_seed": bool(randomize_seed),
"quality": float(quality),
"scheduler": scheduler,
"flow_shift": float(flow_shift),
"frame_multiplier": int(frame_multiplier),
"safe_mode": bool(safe_mode),
"safety_checker": bool(safety_checker),
"play_result_video": bool(play_result_video),
"upscale_model": upscale_model,
"upscale_factor": float(upscale_factor),
}
temp_dir = tempfile.gettempdir()
file_path = os.path.join(temp_dir, f"wan_settings_{uuid.uuid4().hex[:8]}.json")
with open(file_path, "w", encoding="utf-8") as f:
json.dump(settings_dict, f, indent=4, ensure_ascii=False)
return file_path
def extract_json_from_video(file_path):
"""
Extracts embedded JSON metadata from an MP4 video file using ffprobe or ffmpeg.
"""
import imageio_ffmpeg
ffmpeg_exe = imageio_ffmpeg.get_ffmpeg_exe()
ffprobe_exe = os.path.join(os.path.dirname(ffmpeg_exe), "ffprobe")
if not os.path.exists(ffprobe_exe):
ffprobe_exe = "ffprobe"
try:
cmd = [ffprobe_exe, "-v", "quiet", "-print_format", "json", "-show_format", file_path]
res = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, encoding="utf-8", errors="ignore")
if res.returncode == 0 and res.stdout:
probe_data = json.loads(res.stdout)
tags = probe_data.get("format", {}).get("tags", {})
for key, val in tags.items():
if isinstance(val, str) and "prompt" in val:
try:
parsed = json.loads(val)
if isinstance(parsed, dict) and "prompt" in parsed:
return parsed
except Exception:
pass
except Exception:
pass
try:
cmd = [ffmpeg_exe, "-i", file_path]
res = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, encoding="utf-8", errors="ignore")
stderr_text = res.stderr
matches = re.findall(r'\{[^{}]*"prompt"[^{}]*\}', stderr_text, re.DOTALL)
for m in matches:
try:
parsed = json.loads(m)
if isinstance(parsed, dict) and "prompt" in parsed:
return parsed
except Exception:
pass
comment_match = re.search(r'(?:comment|description)\s*:\s*(\{.*\})', stderr_text, re.IGNORECASE)
if comment_match:
try:
parsed = json.loads(comment_match.group(1).strip())
if isinstance(parsed, dict):
return parsed
except Exception:
pass
except Exception:
pass
return None
def parse_settings_file(file_obj):
"""
Reads parameter settings from an uploaded JSON file or MP4 video metadata.
"""
if not file_obj:
return [gr.update() for _ in range(17)]
file_path = file_obj.name if hasattr(file_obj, "name") else str(file_obj)
data = None
try:
with open(file_path, "r", encoding="utf-8") as f:
data = json.load(f)
except Exception:
data = None
if data is None:
data = extract_json_from_video(file_path)
if not isinstance(data, dict):
gr.Warning("No valid generation settings found in file.")
return [gr.update() for _ in range(17)]
gr.Info("Settings loaded successfully!")
return (
data.get("prompt", gr.update()),
data.get("negative_prompt", gr.update()),
data.get("steps", gr.update()),
data.get("duration_seconds", gr.update()),
data.get("guidance_scale", gr.update()),
data.get("guidance_scale_2", gr.update()),
data.get("seed", gr.update()),
data.get("randomize_seed", gr.update()),
data.get("quality", gr.update()),
data.get("scheduler", gr.update()),
data.get("flow_shift", gr.update()),
data.get("frame_multiplier", gr.update()),
data.get("safe_mode", gr.update()),
data.get("safety_checker", gr.update()),
data.get("play_result_video", data.get("video_component", gr.update())),
data.get("upscale_model") if data.get("upscale_model") in UPSCALER_NAMES else (UPSCALER_NAMES[0] if UPSCALER_NAMES else None),
data.get("upscale_factor", gr.update()),
)
def clear_vram():
gc.collect()
torch.cuda.empty_cache()
# RIFE
if not os.path.exists("RIFEv4.26_0921.zip"):
print("Downloading RIFE Model...")
rife_file = "RIFEv4.26_0921.zip"
hf_hub_download(repo_id="r3gm/RIFE", filename="RIFEv4.26_0921.zip", local_dir=".")
subprocess.run(["unzip", "-o", rife_file], check=True)
# sys.path.append(os.getcwd())
from train_log.RIFE_HDv3 import Model
rife_model = Model()
rife_model.load_model("train_log", -1)
rife_model.eval()
UPSCALER_MODELS_CONFIG = {
"2xNomosUni_compact_otf_medium": {
"repo_source": ["Phips/2xNomosUni_compact_otf_medium", "2xNomosUni_compact_otf_medium.safetensors"],
"s_frame": 0.020,
"max_scale": 2.0
},
# "2xHFA2k_LUDVAE_compact": { # no ask
# "repo_source": ["Phips/2xHFA2k_LUDVAE_compact", "2xHFA2k_LUDVAE_compact.safetensors"],
# "s_frame": 0.020,
# "max_scale": 2.0
# },
"R-ESRGAN_AnimeVideo_v3": {
"repo_source": ["iahhnim/ESRGAN_collection", "realesr-animevideov3.pth"],
"s_frame": 0.050,
"max_scale": 4.0
},
# "4x-ClearRealityV1": {
# "repo_source": ["Kim2091/ClearRealityV1", "4x-ClearRealityV1.safetensors"],
# "s_frame": 0.055,
# "max_scale": 4.0
# },
"R-ESRGAN_x4plus_Anime6B": {
"repo_source": ["iahhnim/ESRGAN_collection", "RealESRGAN_x4plus_anime_6B.pth"],
"s_frame": 0.125,
"max_scale": 4.0
},
"4xNomosWebPhoto_RealPLKSR": {
"repo_source": ["Phips/4xNomosWebPhoto_RealPLKSR", "4xNomosWebPhoto_RealPLKSR.safetensors"],
"s_frame": 0.180,
"max_scale": 4.0
},
"4xBHI_realplksr_dysample_otf": {
"repo_source": ["Phips/4xBHI_realplksr_dysample_otf", "4xBHI_realplksr_dysample_otf.safetensors"],
"s_frame": 0.190,
"max_scale": 4.0
},
"4x-AnimeSharp": {
"repo_source": ["Kim2091/AnimeSharp", "4x-AnimeSharp.safetensors"],
"s_frame": 0.290,
"max_scale": 4.0
},
"4x-UltraSharp": {
"repo_source": ["Kim2091/UltraSharp", "4x-UltraSharp.safetensors"],
"s_frame": 0.290,
"max_scale": 4.0
},
# "4xNomosWebPhoto_esrgan": {
# "repo_source": ["Phips/4xNomosWebPhoto_esrgan", "4xNomosWebPhoto_esrgan.safetensors"],
# "s_frame": 0.290,
# "max_scale": 4.0
# },
}
spandrel_loader = ModelLoader()
LOADED_SPANDREL_MODELS = {}
print("Preloading Safetensors Spandrel Upscaler Models at startup...")
for model_key, info in UPSCALER_MODELS_CONFIG.items():
repo_id, local_filename = info["repo_source"]
try:
if not os.path.exists(local_filename):
print(f"Downloading {model_key} ({local_filename}) from {repo_id}...")
cache_file_up = hf_hub_download(repo_id=repo_id, filename=local_filename)
descriptor = spandrel_loader.load_from_file(cache_file_up)
loaded_m = descriptor.model.to(device).half()
loaded_m.eval()
LOADED_SPANDREL_MODELS[model_key] = loaded_m
print(f" Successfully preloaded: {model_key}")
except Exception as e:
print(f" Failed to load {model_key}: {e}")
UPSCALER_NAMES = list(LOADED_SPANDREL_MODELS.keys())
@torch.no_grad()
def upscale_frames_spandrel(frames_list, model_name="4x-UltraSharp", scale_factor=2.0):
"""
Upscales a list of Numpy float32 [0.0, 1.0] frames (H, W, C) using the selected preloaded Spandrel model in FP16.
Enforces native model scaling limits (e.g., 2x max models).
"""
if scale_factor <= 1.0 or not frames_list:
return frames_list, 0, 0.0, 1.0
model_config = UPSCALER_MODELS_CONFIG.get(model_name, {})
max_scale = model_config.get("max_scale", 4.0)
# Enforce native model scale limits
effective_scale = scale_factor
if effective_scale > max_scale:
effective_scale = max_scale
warning_msg = f"Notice: Model '{model_name}' has a maximum native scale of {max_scale}x. Output scale clamped to {effective_scale}x."
print(warning_msg)
gr.Info(warning_msg)
selected_model = LOADED_SPANDREL_MODELS.get(model_name)
start_time = time.time()
num_frames = len(frames_list)
H, W, C = frames_list[0].shape
target_H = int(round(H * effective_scale))
target_W = int(round(W * effective_scale))
# Ensure target dimensions are even (divisible by 2) for standard video codecs
target_H = (target_H // 2) * 2
target_W = (target_W // 2) * 2
upscaled_frames = []
with tqdm(total=num_frames, desc=f"Upscaling with {model_name} ({effective_scale}x)", unit="frame") as pbar:
for i, frame_np in enumerate(frames_list):
# Convert Numpy HWC float32 -> PyTorch Tensor BCHW float16 on GPU
t = torch.from_numpy(frame_np).to(device).permute(2, 0, 1).unsqueeze(0).half()
# Forward pass through selected Spandrel Model
out_t = selected_model(t)
# If target dimension differs from model output, resize using bicubic interpolation
if out_t.shape[2] != target_H or out_t.shape[3] != target_W:
out_t = F.interpolate(out_t, size=(target_H, target_W), mode="bicubic", align_corners=False)
# Clamp and convert back to Numpy HWC float32 [0.0, 1.0]
out_t = out_t.clamp(0.0, 1.0).squeeze(0).permute(1, 2, 0).float().cpu().numpy()
upscaled_frames.append(out_t)
if (i + 1) % 3 == 0:
pbar.update(3)
pbar.update(num_frames % 3)
elapsed_time = time.time() - start_time
torch.cuda.empty_cache()
return upscaled_frames, num_frames, elapsed_time, effective_scale
@torch.no_grad()
def interpolate_bits(frames_np, multiplier=2, scale=1.0):
"""
Interpolation maintaining Numpy Float 0-1 format.
Args:
frames_np: Numpy Array (Time, Height, Width, Channels) - Float32 [0.0, 1.0]
multiplier: int (2, 4, 8)
Returns:
List of Numpy Arrays (Height, Width, Channels) - Float32 [0.0, 1.0]
"""
# Handle input shape
if isinstance(frames_np, list):
# Convert list of arrays to one big array for easier shape handling if needed,
# but here we just grab dims from first frame
T = len(frames_np)
H, W, C = frames_np[0].shape
else:
T, H, W, C = frames_np.shape
# 1. No Interpolation Case
if multiplier < 2:
# Just convert 4D array to list of 3D arrays
if isinstance(frames_np, np.ndarray):
return list(frames_np)
return frames_np
n_interp = multiplier - 1
# Pre-calc padding for RIFE (requires dimensions divisible by 32/scale)
tmp = max(128, int(128 / scale))
ph = ((H - 1) // tmp + 1) * tmp
pw = ((W - 1) // tmp + 1) * tmp
padding = (0, pw - W, 0, ph - H)
# Helper: Numpy (H, W, C) Float -> Tensor (1, C, H, W) Half
def to_tensor(frame_np):
# frame_np is float32 0-1
t = torch.from_numpy(frame_np).to(device)
# HWC -> CHW
t = t.permute(2, 0, 1).unsqueeze(0)
return F.pad(t, padding).half()
# Helper: Tensor (1, C, H, W) Half -> Numpy (H, W, C) Float
def from_tensor(tensor):
# Crop padding
t = tensor[0, :, :H, :W]
# CHW -> HWC
t = t.permute(1, 2, 0)
# Keep as float32, range 0-1
return t.float().cpu().numpy()
def make_inference(I0, I1, n):
if rife_model.version >= 3.9:
res = []
for i in range(n):
res.append(rife_model.inference(I0, I1, (i+1) * 1. / (n+1), scale))
return res
else:
middle = rife_model.inference(I0, I1, scale)
if n == 1:
return [middle]
first_half = make_inference(I0, middle, n=n//2)
second_half = make_inference(middle, I1, n=n//2)
if n % 2:
return [*first_half, middle, *second_half]
else:
return [*first_half, *second_half]
output_frames = []
# Process Frames
# Load first frame into GPU
I1 = to_tensor(frames_np[0])
total_steps = T - 1
with tqdm(total=total_steps, desc="Interpolating", unit="frame") as pbar:
for i in range(total_steps):
I0 = I1
# Add original frame to output
output_frames.append(from_tensor(I0))
# Load next frame
I1 = to_tensor(frames_np[i+1])
# Generate intermediate frames
mid_tensors = make_inference(I0, I1, n_interp)
# Append intermediate frames
for mid in mid_tensors:
output_frames.append(from_tensor(mid))
if (i + 1) % 25 == 0:
pbar.update(25)
pbar.update(total_steps % 25)
# Add the very last frame
output_frames.append(from_tensor(I1))
# Cleanup
del I0, I1, mid_tensors
torch.cuda.empty_cache()
return output_frames
# WAN
ORG_NAME = "TestOrganizationPleaseIgnore"
# MODEL_ID = "Wan-AI/Wan2.2-I2V-A14B-Diffusers"
MODEL_ID = os.getenv("REPO_ID") or random.choice(
list(list_models(author=ORG_NAME, filter='diffusers:WanImageToVideoPipeline'))
).modelId
CACHE_DIR = os.path.expanduser("~/.cache/huggingface/")
LORA_MODELS = [
# {
# "repo_id": "exampleuser/example_lora_1",
# "high_tr": "example_lora_1_high.safetensors",
# "low_tr": "example_lora_1_low.safetensors",
# "high_scale": 0.5,
# "low_scale": 0.5
# },
# {
# "repo_id": "exampleuser/example_lora_2",
# "high_tr": "subfolder/example_lora_2_high.safetensors",
# "low_tr": "subfolder/example_lora_2_low.safetensors",
# "high_scale": 0.4,
# "low_scale": 0.4
# },
]
MAX_DIM = 832
MIN_DIM = 480
SQUARE_DIM = 640
MULTIPLE_OF = 16
MAX_SEED = np.iinfo(np.int32).max
FIXED_FPS = 16
MIN_FRAMES_MODEL = 8
MAX_FRAMES_MODEL = 160
MIN_DURATION = round(MIN_FRAMES_MODEL / FIXED_FPS, 1)
MAX_DURATION = round(MAX_FRAMES_MODEL / FIXED_FPS, 1)
SCHEDULER_MAP = {
"FlowMatchEulerDiscrete": FlowMatchEulerDiscreteScheduler,
"SASolver": SASolverScheduler,
"DEISMultistep": DEISMultistepScheduler,
"DPMSolverMultistepInverse": DPMSolverMultistepInverseScheduler,
"UniPCMultistep": UniPCMultistepScheduler,
"DPMSolverMultistep": DPMSolverMultistepScheduler,
"DPMSolverSinglestep": DPMSolverSinglestepScheduler,
}
pipe = WanImageToVideoPipeline.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16,
).to('cuda')
original_scheduler = copy.deepcopy(pipe.scheduler)
for i, lora in enumerate(LORA_MODELS):
name_high_tr = lora["high_tr"].split(".")[0].split("/")[-1] + "Hh"
name_low_tr = lora["low_tr"].split(".")[0].split("/")[-1] + "Ll"
try:
pipe.load_lora_weights(
lora["repo_id"],
weight_name=lora["high_tr"],
adapter_name=name_high_tr
)
kwargs_lora = {"load_into_transformer_2": True}
pipe.load_lora_weights(
lora["repo_id"],
weight_name=lora["low_tr"],
adapter_name=name_low_tr,
**kwargs_lora
)
pipe.set_adapters([name_high_tr, name_low_tr], adapter_weights=[1.0, 1.0])
pipe.fuse_lora(adapter_names=[name_high_tr], lora_scale=lora["high_scale"], components=["transformer"])
pipe.fuse_lora(adapter_names=[name_low_tr], lora_scale=lora["low_scale"], components=["transformer_2"])
pipe.unload_lora_weights()
print(f"Applied: {lora['high_tr']}, hs={lora['high_scale']}/ls={lora['low_scale']}, {i+1}/{len(LORA_MODELS)}")
except Exception as e:
print("Error:", str(e))
print("Failed LoRA:", name_high_tr)
pipe.unload_lora_weights()
# if os.path.exists(CACHE_DIR):
# shutil.rmtree(CACHE_DIR)
# print("Deleted Hugging Face cache.")
# else:
# print("No hub cache found.")
quantize_(pipe.text_encoder, Int8WeightOnlyConfig())
torch._dynamo.reset()
quantize_(pipe.transformer, Float8DynamicActivationFloat8WeightConfig())
torch._dynamo.reset()
quantize_(pipe.transformer_2, Float8DynamicActivationFloat8WeightConfig())
torch._dynamo.reset()
spaces.aoti_load(
module=pipe.transformer,
repo_id='cbensimon/WanTransformer3DModel-sm120-cu130-raa',
)
spaces.aoti_load(
module=pipe.transformer_2,
repo_id='cbensimon/WanTransformer3DModel-sm120-cu130-raa',
)
# pipe.vae.enable_slicing()
# pipe.vae.enable_tiling()
default_prompt_i2v = "make this image come alive, cinematic motion, smooth animation"
default_negative_prompt = "色调艳丽, 过曝, 静态, 细节模糊不清, 字幕, 风格, 作品, 画作, 画面, 静止, 整体发灰, 最差质量, 低质量, JPEG压缩残留, 丑陋的, 残缺的, 多余的手指, 画得不好的手部, 画得不好的脸部, 畸形的, 毁容的, 形态畸形的肢体, 手指融合, 静止不动的画面, 杂乱的背景, 三条腿, 背景人很多, 倒着走"
def model_title():
repo_name = MODEL_ID.split('/')[-1].replace("_", " ")
url = f"https://huggingface.co/{MODEL_ID}"
return f"### This space is currently running [{repo_name}]({url}) 🐢"
def resize_image(image: Image.Image) -> Image.Image:
width, height = image.size
if width == height:
return image.resize((SQUARE_DIM, SQUARE_DIM), Image.LANCZOS)
aspect_ratio = width / height
MAX_ASPECT_RATIO = MAX_DIM / MIN_DIM
MIN_ASPECT_RATIO = MIN_DIM / MAX_DIM
image_to_resize = image
if aspect_ratio > MAX_ASPECT_RATIO:
target_w, target_h = MAX_DIM, MIN_DIM
crop_width = int(round(height * MAX_ASPECT_RATIO))
left = (width - crop_width) // 2
image_to_resize = image.crop((left, 0, left + crop_width, height))
elif aspect_ratio < MIN_ASPECT_RATIO:
target_w, target_h = MIN_DIM, MAX_DIM
crop_height = int(round(width / MIN_ASPECT_RATIO))
top = (height - crop_height) // 2
image_to_resize = image.crop((0, top, width, top + crop_height))
else:
if width > height:
target_w = MAX_DIM
target_h = int(round(target_w / aspect_ratio))
else:
target_h = MAX_DIM
target_w = int(round(target_h * aspect_ratio))
final_w = round(target_w / MULTIPLE_OF) * MULTIPLE_OF
final_h = round(target_h / MULTIPLE_OF) * MULTIPLE_OF
final_w = max(MIN_DIM, min(MAX_DIM, final_w))
final_h = max(MIN_DIM, min(MAX_DIM, final_h))
return image_to_resize.resize((final_w, final_h), Image.LANCZOS)
def resize_and_crop_to_match(target_image, reference_image):
ref_width, ref_height = reference_image.size
target_width, target_height = target_image.size
scale = max(ref_width / target_width, ref_height / target_height)
new_width, new_height = int(target_width * scale), int(target_height * scale)
resized = target_image.resize((new_width, new_height), Image.Resampling.LANCZOS)
left, top = (new_width - ref_width) // 2, (new_height - ref_height) // 2
return resized.crop((left, top, left + ref_width, top + ref_height))
def get_num_frames(duration_seconds: float):
return 1 + int(np.clip(
int(round(duration_seconds * FIXED_FPS)),
MIN_FRAMES_MODEL,
MAX_FRAMES_MODEL,
))
def get_inference_duration(
resized_image,
processed_last_image,
prompt,
steps,
negative_prompt,
num_frames,
guidance_scale,
guidance_scale_2,
current_seed,
randomize_seed,
scheduler_name,
flow_shift,
frame_multiplier,
quality,
duration_seconds,
upscale_model,
upscale_factor,
safe_mode,
enable_safety_checker,
video_component,
progress
):
BASE_FRAMES_HEIGHT_WIDTH = 161 * 832 * 624
BASE_STEP_DURATION = 21.
width, height = resized_image.size
factor = num_frames * width * height / BASE_FRAMES_HEIGHT_WIDTH
step_duration = BASE_STEP_DURATION * factor ** 1.7
gen_time = int(steps) * step_duration
if guidance_scale > 1:
gen_time = gen_time * 2.4
frame_factor = frame_multiplier // FIXED_FPS
total_out_frames = (num_frames * frame_factor) if frame_factor > 1 else num_frames
if frame_factor > 1:
inter_time = ((total_out_frames - num_frames) * 0.02)
gen_time += inter_time
if float(upscale_factor) > 1.0:
ups_config = UPSCALER_MODELS_CONFIG.get(upscale_model, {"s_frame": 0.0, "max_scale": 1.0})
gen_time += total_out_frames * (ups_config["s_frame"] + 0.002)
total_time = 12 + gen_time
if safe_mode:
total_time = total_time * 1.25
# print(total_time)
return total_time
@spaces.GPU(duration=get_inference_duration, size='xlarge')
def run_inference(
resized_image,
processed_last_image,
prompt,
steps,
negative_prompt,
num_frames,
guidance_scale,
guidance_scale_2,
current_seed,
randomize_seed,
scheduler_name,
flow_shift,
frame_multiplier,
quality,
duration_seconds,
upscale_model,
upscale_factor,
safe_mode,
enable_safety_checker,
video_component,
progress=gr.Progress(track_tqdm=True),
):
scheduler_class = SCHEDULER_MAP.get(scheduler_name)
if scheduler_class.__name__ != pipe.scheduler.config._class_name or flow_shift != pipe.scheduler.config.get("flow_shift", "shift"):
config = copy.deepcopy(original_scheduler.config)
if scheduler_class == FlowMatchEulerDiscreteScheduler:
config['shift'] = flow_shift
else:
config['flow_shift'] = flow_shift
pipe.scheduler = scheduler_class.from_config(config)
clear_vram()
task_name = str(uuid.uuid4())[:8]
if enable_safety_checker:
if check_nsfw(resized_image, prompt) or (processed_last_image is not None and check_nsfw(processed_last_image, prompt)):
return None, task_name, True
print(f"Generating {num_frames} frames, task: {task_name}, {duration_seconds}, {resized_image.size}")
start = time.time()
result = pipe(
image=resized_image,
last_image=processed_last_image,
prompt=prompt,
negative_prompt=negative_prompt,
height=resized_image.height,
width=resized_image.width,
num_frames=num_frames,
guidance_scale=float(guidance_scale),
guidance_scale_2=float(guidance_scale_2),
num_inference_steps=int(steps),
generator=torch.Generator(device="cuda").manual_seed(current_seed),
output_type="np"
)
print("gen time passed:", time.time() - start)
raw_frames_np = result.frames[0] # Returns (T, H, W, C) float32
pipe.scheduler = original_scheduler
is_nsfw = False
if enable_safety_checker:
if processed_last_image is None:
is_nsfw = check_nsfw(raw_frames_np[-1])
if is_nsfw:
return None, task_name, True
frame_factor = frame_multiplier // FIXED_FPS
if frame_factor > 1:
start = time.time()
print(f"Processing frames (RIFE Multiplier: {frame_factor}x)...")
rife_model.device()
rife_model.flownet = rife_model.flownet.half()
final_frames = interpolate_bits(raw_frames_np, multiplier=int(frame_factor))
print("Interpolation time passed:", time.time() - start)
else:
final_frames = list(raw_frames_np)
# --- SPANDREL POST-PROCESSING FRAME UPSCALE ---
if float(upscale_factor) > 1.0:
print(f"Upscaling frames using Spandrel ({upscale_model}, Scale: {upscale_factor}x)...")
final_frames, up_count, up_time, applied_scale = upscale_frames_spandrel(
final_frames, model_name=upscale_model, scale_factor=float(upscale_factor)
)
per_frame_time = up_time / max(1, up_count)
mini_report = f"Upscaling Report: {up_count} frames upscaled to {applied_scale}x using '{upscale_model}' in {up_time:.2f}s ({per_frame_time:.3f}s/frame)"
print(mini_report)
gr.Info(mini_report)
# ---------------------------------------------
final_fps = FIXED_FPS * int(frame_factor)
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmpfile:
video_path = tmpfile.name
metadata_dict = {
"prompt": prompt,
"negative_prompt": negative_prompt,
"steps": int(steps),
"duration_seconds": float(duration_seconds),
"guidance_scale": float(guidance_scale),
"guidance_scale_2": float(guidance_scale_2),
"seed": int(current_seed),
"randomize_seed": bool(randomize_seed),
"quality": float(quality),
"scheduler": scheduler_name,
"flow_shift": float(flow_shift),
"frame_multiplier": int(frame_multiplier),
"upscale_model": upscale_model,
"upscale_factor": float(upscale_factor),
"safe_mode": bool(safe_mode),
"safety_checker": bool(enable_safety_checker),
"play_result_video": bool(video_component),
}
start = time.time()
export_to_video(final_frames, video_path, fps=final_fps, quality=quality, metadata=metadata_dict)
print(f"Export time passed, {final_fps} FPS:", time.time() - start)
return video_path, task_name, False
def generate_video(
input_image,
last_image,
prompt,
steps=4,
negative_prompt=default_negative_prompt,
duration_seconds=MAX_DURATION,
guidance_scale=1,
guidance_scale_2=1,
seed=42,
randomize_seed=False,
quality=5,
scheduler="UniPCMultistep",
flow_shift=6.0,
frame_multiplier=16,
upscale_model="4x-UltraSharp",
upscale_factor=1.0,
video_component=True,
safe_mode=False,
enable_safety_checker=True,
progress=gr.Progress(track_tqdm=True),
):
"""
Generate a video from an input image using the Wan 2.2 14B I2V model with Lightning LoRA.
This function takes an input image and generates a video animation based on the provided
prompt and parameters. It uses an FP8 qunatized Wan 2.2 14B Image-to-Video model in with Lightning LoRA
for fast generation in 4-8 steps.
Args:
input_image (PIL.Image): The input image to animate. Will be resized to target dimensions.
last_image (PIL.Image, optional): The optional last image for the video.
prompt (str): Text prompt describing the desired animation or motion.
steps (int, optional): Number of inference steps. More steps = higher quality but slower.
Defaults to 4. Range: 1-30.
negative_prompt (str, optional): Negative prompt to avoid unwanted elements.
Defaults to default_negative_prompt (contains unwanted visual artifacts).
duration_seconds (float, optional): Duration of the generated video in seconds.
Defaults to 2. Clamped between MIN_FRAMES_MODEL/FIXED_FPS and MAX_FRAMES_MODEL/FIXED_FPS.
guidance_scale (float, optional): Controls adherence to the prompt. Higher values = more adherence.
Defaults to 1.0. Range: 0.0-20.0.
guidance_scale_2 (float, optional): Controls adherence to the prompt. Higher values = more adherence.
Defaults to 1.0. Range: 0.0-20.0.
seed (int, optional): Random seed for reproducible results. Defaults to 42.
Range: 0 to MAX_SEED (2147483647).
randomize_seed (bool, optional): Whether to use a random seed instead of the provided seed.
Defaults to False.
quality (float, optional): Video output quality. Default is 5. Uses variable bit rate.
Highest quality is 10, lowest is 1.
scheduler (str, optional): The name of the scheduler to use for inference. Defaults to "UniPCMultistep".
flow_shift (float, optional): The flow shift value for compatible schedulers. Defaults to 6.0.
frame_multiplier (int, optional): The int value for fps enhancer
upscale_model (str, optional): Selected Spandrel upscaler model architecture.
upscale_factor (float, optional): Post-processing scale factor for Spandrel upscaler (1.0 to 4.0).
video_component(bool, optional): Show video player in output.
Defaults to True.
progress (gr.Progress, optional): Gradio progress tracker. Defaults to gr.Progress(track_tqdm=True).
Returns:
tuple: A tuple containing:
- video_path (str): Path for the video component.
- video_path (str): Path for the file download component. Attempt to avoid reconversion in video component.
- current_seed (int): The seed used for generation.
Raises:
gr.Error: If input_image is None (no image uploaded).
Note:
- Frame count is calculated as duration_seconds * FIXED_FPS (24)
- Output dimensions are adjusted to be multiples of MOD_VALUE (32)
- The function uses GPU acceleration via the @spaces.GPU decorator
- Generation time varies based on steps and duration (see get_duration function)
"""
if input_image is None:
raise gr.Error("Please upload an input image.")
current_seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed)
if pre_gpu_safety_check(input_image, last_image, prompt):
gr.Warning("Generation blocked: This request was flagged by a content filter and wasn't run.")
return gr.update(), gr.update(), current_seed
num_frames = get_num_frames(duration_seconds)
resized_image = resize_image(input_image)
processed_last_image = None
if last_image:
processed_last_image = resize_and_crop_to_match(last_image, resized_image)
video_path, task_n, is_nsfw = run_inference(
resized_image,
processed_last_image,
prompt,
steps,
negative_prompt,
num_frames,
guidance_scale,
guidance_scale_2,
current_seed,
randomize_seed,
scheduler,
flow_shift,
frame_multiplier,
quality,
duration_seconds,
upscale_model,
upscale_factor,
safe_mode,
enable_safety_checker,
video_component,
progress,
)
if is_nsfw:
gr.Warning("Generation blocked by guardrails: The resulting video may contain sensitive or explicit content.")
return gr.update(), gr.update(), current_seed
print(f"GPU complete: {task_n}")
return (video_path if video_component else None), [video_path], current_seed
with gr.Blocks(delete_cache=(1200, 7200), fill_width=True) as demo:
gr.Markdown(model_title())
gr.Markdown("Run Wan 2.2 in just 4-8 steps, fp8 quantization & AoT compilation - compatible with 🧨 diffusers and ZeroGPU")
with gr.Row(elem_id="main-container"):
with gr.Column():
input_image_component = gr.Image(type="pil", label="Input Image", sources=["upload", "clipboard"], buttons=["fullscreen"])
prompt_input = gr.Textbox(label="Prompt", value=default_prompt_i2v)
with gr.Row():
duration_seconds_input = gr.Slider(minimum=MIN_DURATION, maximum=MAX_DURATION, step=0.1, value=3.5, label="Duration (seconds)")
frame_multi = gr.Dropdown(
choices=[FIXED_FPS, FIXED_FPS*2, FIXED_FPS*4, FIXED_FPS*8],
value=FIXED_FPS,
label="FPS"
)
with gr.Accordion("Advanced Settings", open=False):
with gr.Row(elem_classes=["compact-btn-row"]):
load_settings_btn = gr.UploadButton(
"Load Settings (JSON / MP4)",
file_types=[".json", ".mp4", "video/*"],
file_count="single",
size="sm",
elem_classes=["compact-btn"]
)
download_json_btn = gr.DownloadButton("Download Settings JSON", size="sm", elem_classes=["compact-btn"])
last_image_component = gr.Image(type="pil", label="Last Image (Optional)", sources=["upload", "clipboard"], buttons=["fullscreen"])
negative_prompt_input = gr.Textbox(label="Negative Prompt", value=default_negative_prompt, info="Used if Guidance Scale > 1.", lines=3)
steps_slider = gr.Slider(minimum=1, maximum=12, step=1, value=6, label="Inference Steps")
quality_slider = gr.Slider(minimum=1, maximum=10, step=1, value=6, label="Video Encoding Quality (Bitrate)")
# --- SPANDREL SAFETENSORS FRAME UPSCALER CONTROLS ---
upscaler_dropdown = gr.Dropdown(
choices=UPSCALER_NAMES,
value=UPSCALER_NAMES[0] if UPSCALER_NAMES else None,
label="Upscaler Model",
info="Select a upscaler model."
)
upscale_factor_slider = gr.Slider(
minimum=1.0, maximum=4.0, step=0.5, value=1.0,
label="Frame Upscale Factor",
info="Upscales output frames post-interpolation using selected model (1.0 = Off)."
)
# ----------------------------------------------------
scheduler_dropdown = gr.Dropdown(
label="Scheduler",
choices=list(SCHEDULER_MAP.keys()),
value="UniPCMultistep",
info="Select a custom scheduler."
)
flow_shift_slider = gr.Slider(minimum=0.5, maximum=15.0, step=0.1, value=3.0, label="Flow Shift")
seed_input = gr.Slider(label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=42, interactive=True)
randomize_seed_checkbox = gr.Checkbox(label="Randomize seed", value=True, interactive=True)
guidance_scale_input = gr.Slider(minimum=0.0, maximum=10.0, step=0.5, value=1, label="Guidance Scale - high noise stage", info="Values above 1 increase GPU usage and may take longer to process.")
guidance_scale_2_input = gr.Slider(minimum=0.0, maximum=10.0, step=0.5, value=1, label="Guidance Scale 2 - low noise stage")
safe_mode_checkbox = gr.Checkbox(
label="🛠️ Safe Mode",
value=True,
info="Requests 25% extra processing time to try to prevent unfinished tasks when the server is busy."
)
safety_checker_input = gr.Checkbox(label="Enable Safety Filter", value=True, info="Prevents unrequested sensitive or explicit content.")
play_result_video = gr.Checkbox(label="Display result", value=True, interactive=True)
with gr.Accordion("🎬 Video Merger", open=False):
send_to_merge_btn = gr.Button("➕ Add Video (or Part 1) to Merger", size="sm", elem_classes=["compact-btn"])
merge_files_input = gr.File(label="Videos to Merge", file_count="multiple", file_types=[".mp4", "video/*"])
merge_btn = gr.Button("🔗 Merge Videos", variant="secondary")
merged_file_output = gr.File(label="Download Merged Video")
gr.Markdown(f"[ZeroGPU help, tips and troubleshooting](https://huggingface.co/datasets/{ORG_NAME}/help/blob/main/gpu_help.md)")
gr.Markdown( # TestOrganizationPleaseIgnore/wamu-tools
"To use a different model, **duplicate this Space** first, then change the `REPO_ID` environment variable. "
"[See compatible models here](https://huggingface.co/models?other=diffusers:WanImageToVideoPipeline&sort=trending&search=WAN2.2_I2V_LIGHTNING)."
)
generate_button = gr.Button("Generate Video", variant="primary", elem_classes=["generate-btn"])
with gr.Column():
# ASSIGNED elem_id="generated-video" so JS can find it
video_output = gr.Video(label="Generated Video", autoplay=True, sources=["upload"], buttons=["download", "share"], interactive=False, elem_id="generated-video")
# --- Frame Grabbing UI ---
with gr.Row(elem_classes=["compact-btn-row"]):
step_back_btn = gr.Button("◀", size="sm", elem_classes=["compact-btn"])
grab_frame_btn = gr.Button("📸 Use Current Frame as Input", size="sm", elem_classes=["compact-btn"])
split_btn = gr.Button("✂️ Split", size="sm", elem_classes=["compact-btn"])
step_fwd_btn = gr.Button("▶", size="sm", elem_classes=["compact-btn"])
frame_timestamp_box = gr.Number(value=0, label="Frame Timestamp", visible=False)
split_timestamp_box = gr.Number(value=0, label="Split Timestamp", visible=False)
# -------------------------
file_output = gr.File(label="Download Video", file_count="multiple", interactive=False)
gr.Examples(
examples=[
[
"wan_i2v_input.JPG",
"POV selfie video, white cat with sunglasses standing on surfboard, relaxed smile, tropical beach behind (clear water, green hills, blue sky with clouds). Surfboard tips, cat falls into ocean, camera plunges underwater with bubbles and sunlight beams. Brief underwater view of cat's face, then cat resurfaces, still filming selfie, playful summer vacation mood."
]
],
inputs=[input_image_component, prompt_input],
cache_examples=False,
)
ui_inputs = [
input_image_component, last_image_component, prompt_input, steps_slider,
negative_prompt_input, duration_seconds_input,
guidance_scale_input, guidance_scale_2_input, seed_input, randomize_seed_checkbox,
quality_slider, scheduler_dropdown, flow_shift_slider, frame_multi,
upscaler_dropdown, upscale_factor_slider, play_result_video, safe_mode_checkbox, safety_checker_input
]
generate_button.click(
fn=generate_video,
inputs=ui_inputs,
outputs=[video_output, file_output, seed_input]
)
download_json_btn.click(
fn=export_settings_json,
inputs=[
prompt_input, negative_prompt_input, steps_slider, duration_seconds_input,
guidance_scale_input, guidance_scale_2_input, seed_input, randomize_seed_checkbox, quality_slider,
scheduler_dropdown, flow_shift_slider, frame_multi, safe_mode_checkbox, safety_checker_input,
play_result_video, upscaler_dropdown, upscale_factor_slider
],
outputs=[download_json_btn],
api_visibility="undocumented",
)
load_settings_btn.upload(
fn=parse_settings_file,
inputs=[load_settings_btn],
outputs=[
prompt_input, negative_prompt_input, steps_slider, duration_seconds_input,
guidance_scale_input, guidance_scale_2_input, seed_input, randomize_seed_checkbox,
quality_slider, scheduler_dropdown, flow_shift_slider, frame_multi,
safe_mode_checkbox, safety_checker_input, play_result_video,
upscaler_dropdown, upscale_factor_slider
],
api_visibility="undocumented",
)
# --- Frame Grabbing Events ---
step_back_btn.click(
fn=None,
inputs=None,
outputs=None,
js=step_back_js
)
step_fwd_btn.click(
fn=None,
inputs=None,
outputs=None,
js=step_fwd_js
)
grab_frame_btn.click(
fn=extract_frame,
inputs=[video_output, frame_timestamp_box],
outputs=[input_image_component],
js="""
(video, ts) => {
const videoEl = document.querySelector('#generated-video video');
let currentTime = 0;
if (videoEl) {
videoEl.pause();
currentTime = videoEl.currentTime || 0;
}
return [video, currentTime];
}
""",
api_visibility="private",
)
split_btn.click(
fn=split_video_frame_level,
inputs=[video_output, split_timestamp_box],
outputs=[file_output],
js="""
(video, ts) => {
const videoEl = document.querySelector('#generated-video video');
let currentTime = 0;
if (videoEl) {
videoEl.pause();
currentTime = videoEl.currentTime || 0;
}
return [video, currentTime];
}
""",
api_visibility="private",
)
send_to_merge_btn.click(
fn=send_first_file_to_merger,
inputs=[file_output, merge_files_input],
outputs=[merge_files_input],
api_visibility="private",
)
merge_btn.click(
fn=merge_videos,
inputs=[merge_files_input],
outputs=[merged_file_output],
api_visibility="private",
)
if __name__ == "__main__":
demo.queue().launch(
css=css,
mcp_server=True,
show_error=True,
theme=gr.themes.Monochrome(),
ssr_mode=True,
max_file_size="70mb",
)