wan22-extend / app.py
Simzy's picture
Restore pre-Runpod ZeroGPU Extend (match extend-current-20260923 / GitHub main)
47644d5 verified
Raw History Blame Contribute Delete
27.6 kB
"""
Wan 2.2 Image-to-Video with Extend chaining.
Calls a public ZeroGPU Wan 2.2 I2V Space per segment, extracts the last frame,
and concatenates clips for longer videos (~10–20s).
"""
from __future__ import annotations
import os
import shutil
import subprocess
import tempfile
import time
import uuid
from pathlib import Path
from typing import Any
import cv2
import gradio as gr
import numpy as np
from gradio_client import Client, handle_file
from PIL import Image, ImageOps
# Upstream Wan 2.2 I2V (ZeroGPU) — same Space Command Center uses
UPSTREAM_SPACE = os.environ.get("WAN_UPSTREAM_SPACE", "kulkas2pintu/wan222")
UPSTREAM_API = "generate_video"
DEFAULT_PROMPT = "make this image come alive, cinematic motion, smooth animation"
DEFAULT_NEGATIVE = (
"色调艳丽, 过曝, 静态, 细节模糊不清, 字幕, 风格, 作品, 画作, 画面, 静止, "
"整体发灰, 最差质量, 低质量, JPEG压缩残留, 丑陋的, 残缺的, 多余的手指, "
"画得不好的手部, 画得不好的脸部, 畸形的, 毁容的, 形态畸形的肢体, 手指融合, "
"静止不动的画面, 杂乱的背景, 三条腿, 背景人很多, 倒着走"
)
EXTEND_START_LAST_FRAME = "Last frame from video"
EXTEND_START_UPLOAD = "Upload custom image"
WORK = Path(tempfile.gettempdir()) / "wan22_extend"
WORK.mkdir(parents=True, exist_ok=True)
_client: Client | None = None
def get_client() -> Client:
global _client
if _client is None:
token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
kwargs: dict[str, Any] = {}
if token:
kwargs["hf_token"] = token
_client = Client(UPSTREAM_SPACE, **kwargs)
return _client
def _session_dir() -> Path:
d = WORK / uuid.uuid4().hex[:12]
d.mkdir(parents=True, exist_ok=True)
return d
def _as_image_source(img, _depth: int = 0):
"""Unwrap Gradio file-data / (image, caption) values to a PIL-loadable source."""
if img is None or _depth > 4:
return img
if isinstance(img, dict):
inner = img.get("path") or img.get("name") or img.get("image") or img.get("url")
return _as_image_source(inner, _depth + 1)
if isinstance(img, (list, tuple)) and img:
return _as_image_source(img[0], _depth + 1)
return img
def load_pil_image(img: Image.Image | np.ndarray | str | Path | dict | None) -> Image.Image | None:
"""Load an upload as RGB, honoring iPhone EXIF orientation."""
img = _as_image_source(img)
if img is None:
return None
if isinstance(img, Image.Image):
pil = img
elif isinstance(img, np.ndarray):
arr = img
if arr.ndim == 2:
arr = np.stack([arr, arr, arr], axis=-1)
if arr.shape[-1] == 4:
arr = arr[..., :3]
pil = Image.fromarray(arr.astype("uint8"))
elif isinstance(img, (str, Path)):
src = Path(img)
if not src.exists() or not src.is_file():
return None
try:
with Image.open(src) as opened:
pil = opened.copy()
except Exception:
return None
else:
return None
try:
pil = ImageOps.exif_transpose(pil)
except Exception:
pass
return pil.convert("RGB")
def save_image(img: Image.Image | np.ndarray | str | Path | dict | None, dest: Path) -> Path | None:
pil = load_pil_image(img)
if pil is None:
return None
out = dest.with_suffix(".png")
pil.save(out)
return out
def _even(n: int) -> int:
n = int(n)
if n < 2:
return 2
return n - (n % 2)
def probe_video_size(video_path: str | Path) -> tuple[int, int]:
"""Even pixel size of the first video stream (the size players lock onto)."""
video_path = str(video_path)
try:
out = subprocess.check_output(
[
"ffprobe",
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height",
"-of",
"csv=p=0:s=x",
video_path,
],
text=True,
stderr=subprocess.DEVNULL,
).strip()
w_s, h_s = out.split("x", 1)
w, h = _even(int(w_s)), _even(int(h_s))
if w >= 2 and h >= 2:
return w, h
except Exception:
pass
cap = cv2.VideoCapture(video_path)
w = _even(int(cap.get(cv2.CAP_PROP_FRAME_WIDTH) or 0))
h = _even(int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT) or 0))
cap.release()
if w < 2 or h < 2:
raise RuntimeError(f"Could not read video size: {video_path}")
return w, h
def cover_resize(image: Image.Image, width: int, height: int) -> Image.Image:
"""Scale-and-center-crop so `image` fills width×height without letterboxing."""
width, height = _even(width), _even(height)
image = image.convert("RGB")
src_w, src_h = image.size
if src_w < 1 or src_h < 1:
raise ValueError("Custom image has no pixels.")
if src_w == width and src_h == height:
return image
scale = max(width / src_w, height / src_h)
new_w = max(1, int(round(src_w * scale)))
new_h = max(1, int(round(src_h * scale)))
resized = image.resize((new_w, new_h), Image.Resampling.LANCZOS)
left = max(0, (new_w - width) // 2)
top = max(0, (new_h - height) // 2)
cropped = resized.crop((left, top, left + width, top + height))
if cropped.size != (width, height):
cropped = cropped.resize((width, height), Image.Resampling.LANCZOS)
return cropped
def reference_video_path(state: dict) -> Path | None:
"""Clip whose geometry later segments must match. Prefer the original segment."""
segs = state.get("segments") or []
if segs:
first = Path(str(segs[0]))
if first.exists():
return first
video = state.get("video")
if video:
current = Path(str(video))
if current.exists():
return current
return None
def extract_last_frame(video_path: str | Path, out_path: Path) -> Path:
video_path = str(video_path)
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise RuntimeError(f"Could not open video: {video_path}")
total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
frame = None
if total > 0:
cap.set(cv2.CAP_PROP_POS_FRAMES, max(0, total - 1))
ok, frame = cap.read()
if not ok or frame is None:
cap.set(cv2.CAP_PROP_POS_FRAMES, 0)
while True:
ok, f = cap.read()
if not ok:
break
frame = f
else:
while True:
ok, f = cap.read()
if not ok:
break
frame = f
cap.release()
if frame is None:
raise RuntimeError("No frames found in video.")
out_path = out_path.with_suffix(".png")
cv2.imwrite(str(out_path), frame)
return out_path
def video_duration_seconds(video_path: str | Path) -> float:
video_path = str(video_path)
try:
out = subprocess.check_output(
[
"ffprobe",
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
video_path,
],
text=True,
stderr=subprocess.DEVNULL,
).strip()
return float(out)
except Exception:
cap = cv2.VideoCapture(video_path)
fps = cap.get(cv2.CAP_PROP_FPS) or 16.0
frames = cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0
cap.release()
if fps > 0 and frames > 0:
return frames / fps
return 0.0
# Wan I2V on kulkas2pintu/wan222 exports at FIXED_FPS (16) unless fluidity
# interpolation is requested. Concat has always retimed segments to 16fps.
OUTPUT_FPS = 16
def concat_videos(paths: list[Path], out_path: Path) -> Path:
if not paths:
raise RuntimeError("No videos to concatenate.")
if len(paths) == 1:
shutil.copy2(paths[0], out_path)
return out_path
# iOS/Safari lock onto the first sample's resolution, pixel format, and
# SAR. The old scale=trunc(iw/2)*2 kept each clip's own size, so a custom
# extend still (different aspect → different Wan output size) was stream-
# copied in as a mid-file geometry change. Players then hold the last
# decodable frame — a freeze exactly at the custom-frame join.
width, height = probe_video_size(paths[0])
fps = OUTPUT_FPS
# Cover-crop into the first clip's canvas. `pad` rejects inputs that are
# already larger on one side (portrait → landscape), which is the custom-
# image case. Crop with a clamped window, then scale, so a 1px rounding
# shortfall cannot fail the filter.
crop_w = f"if(lt(iw\\,{width})\\,iw\\,{width})"
crop_h = f"if(lt(ih\\,{height})\\,ih\\,{height})"
vf = (
f"scale={width}:{height}:force_original_aspect_ratio=increase,"
f"crop=w={crop_w}:h={crop_h},"
f"scale={width}:{height},"
f"setsar=1,fps={fps},format=yuv420p"
)
list_file = out_path.parent / "concat_list.txt"
normalized: list[Path] = []
for i, p in enumerate(paths):
norm = out_path.parent / f"norm_{i:03d}.mp4"
subprocess.run(
[
"ffmpeg",
"-y",
"-i",
str(p),
"-vf",
vf,
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-an",
"-vsync",
"cfr",
"-r",
str(fps),
str(norm),
],
check=True,
capture_output=True,
)
normalized.append(norm)
with list_file.open("w") as f:
for p in normalized:
escaped = str(p.resolve()).replace("'", r"'\''")
f.write(f"file '{escaped}'\n")
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
str(list_file),
"-c",
"copy",
"-movflags",
"+faststart",
str(out_path),
],
check=True,
capture_output=True,
)
return out_path
def call_wan_i2v(
image_path: Path,
prompt: str,
duration: float,
steps: int,
negative: str,
seed: int,
randomize: bool,
quality: int,
fps: int,
safe_mode: bool,
progress: gr.Progress | None = None,
) -> Path:
client = get_client()
if progress:
progress(0.1, desc=f"Calling {UPSTREAM_SPACE} (ZeroGPU queue)…")
# Match generate_video argument order from kulkas2pintu/wan222 /config
result = client.predict(
handle_file(str(image_path)), # input image
None, # last image (optional end-frame; we chain via last-frame start instead)
prompt or DEFAULT_PROMPT,
int(steps),
negative or DEFAULT_NEGATIVE,
float(duration),
1.0, # guidance high
1.0, # guidance low
int(seed),
bool(randomize),
int(quality),
"UniPCMultistep",
3.0, # flow shift
int(fps),
True, # display result
bool(safe_mode),
api_name=f"/{UPSTREAM_API}",
)
# Returns (video, download_file, seed) typically
video_ref = None
if isinstance(result, (list, tuple)):
video_ref = result[0]
else:
video_ref = result
if isinstance(video_ref, dict):
video_path = video_ref.get("video") or video_ref.get("path") or video_ref.get("url")
if isinstance(video_path, dict):
video_path = video_path.get("path") or video_path.get("url")
else:
video_path = video_ref
if not video_path or not Path(str(video_path)).exists():
raise RuntimeError(
f"Upstream did not return a local video path. Got: {type(result)} {str(result)[:300]}"
)
dest = image_path.parent / f"seg_{uuid.uuid4().hex[:8]}.mp4"
shutil.copy2(str(video_path), dest)
return dest
def status_line(segments: list[str], duration_est: float) -> str:
n = len(segments)
return (
f"**Segments:** {n} · **Duration ≈ {duration_est:.1f}s** · "
f"Upstream: `{UPSTREAM_SPACE}` (ZeroGPU — free, queued/quota-limited)"
)
def estimate_total(seg_duration: float, num_segments: int) -> float:
return float(seg_duration) * int(num_segments)
def resolve_extend_prompt(extend_prompt: str | None, generate_prompt: str | None, state: dict) -> str:
"""Prefer dedicated Extend prompt; else Generate prompt / saved state / default."""
if extend_prompt and str(extend_prompt).strip():
return str(extend_prompt).strip()
if generate_prompt and str(generate_prompt).strip():
return str(generate_prompt).strip()
return state.get("prompt") or DEFAULT_PROMPT
def resolve_extend_start_image(
extend_start: str,
custom_image,
state: dict,
sess: Path,
) -> Path:
"""Pick I2V start image for an Extend segment."""
mode = (extend_start or EXTEND_START_LAST_FRAME).strip()
if mode == EXTEND_START_UPLOAD:
if custom_image is None:
raise gr.Error("Upload a custom extend start image, or switch to “Last frame from video”.")
pil = load_pil_image(custom_image)
if pil is None:
raise gr.Error("Could not read the custom extend start image.")
# Match the clip Wan already produced. Upstream resize_image keeps
# aspect ratio, so a differently shaped still (typical iPhone photo)
# comes back at another resolution and the join freezes on playback.
ref = reference_video_path(state)
if ref is not None:
try:
w, h = probe_video_size(ref)
pil = cover_resize(pil, w, h)
except Exception as e:
raise gr.Error(
f"Could not match the custom image to the current video size: {e}"
) from e
saved = save_image(pil, sess / f"extend_custom_{uuid.uuid4().hex[:8]}")
if saved is None:
raise gr.Error("Could not read the custom extend start image.")
return saved
# Default: last frame from current video
last = Path(state["last_frame"]) if state.get("last_frame") else None
if last is None or not last.exists():
if not state.get("video"):
raise gr.Error("Generate a clip first, then use Extend.")
last = extract_last_frame(state["video"], sess / "last_frame")
state["last_frame"] = str(last)
return last
def do_generate(
image,
prompt,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state: dict,
progress=gr.Progress(track_tqdm=False),
):
if image is None:
raise gr.Error("Upload a still image first.")
sess = Path(state.get("dir") or str(_session_dir()))
sess.mkdir(parents=True, exist_ok=True)
state["dir"] = str(sess)
img_path = save_image(image, sess / "input")
if img_path is None:
raise gr.Error("Could not read the uploaded image.")
progress(0.05, desc="Starting generation…")
try:
clip = call_wan_i2v(
img_path,
prompt,
float(seg_duration),
int(steps),
negative,
int(seed),
bool(randomize),
int(quality),
int(fps),
bool(safe_mode),
progress,
)
except Exception as e:
raise gr.Error(
f"Generation failed (ZeroGPU queue/quota or upstream error): {e}"
) from e
out = sess / "current.mp4"
shutil.copy2(clip, out)
last = extract_last_frame(out, sess / "last_frame")
dur = video_duration_seconds(out)
state["segments"] = [str(clip)]
state["video"] = str(out)
state["last_frame"] = str(last)
state["prompt"] = prompt
progress(1.0, desc="Done")
return (
str(out),
str(out),
status_line(state["segments"], dur),
str(last),
state,
)
def do_extend(
prompt,
extend_prompt,
extend_start,
custom_extend_image,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state: dict,
progress=gr.Progress(track_tqdm=False),
):
if not state or not state.get("video"):
raise gr.Error("Generate a clip first, then use Extend.")
if len(state.get("segments", [])) >= 6:
raise gr.Error("Segment cap reached (6). Reset and start a new chain.")
sess = Path(state["dir"])
start_img = resolve_extend_start_image(extend_start, custom_extend_image, state, sess)
use_prompt = resolve_extend_prompt(extend_prompt, prompt, state)
mode_label = (extend_start or EXTEND_START_LAST_FRAME).strip()
progress(0.05, desc=f"Extending ({mode_label})…")
try:
clip = call_wan_i2v(
start_img,
use_prompt,
float(seg_duration),
int(steps),
negative,
int(seed),
bool(randomize),
int(quality),
int(fps),
bool(safe_mode),
progress,
)
except Exception as e:
raise gr.Error(f"Extend failed: {e}") from e
segs = [Path(p) for p in state["segments"]] + [clip]
out = sess / f"current_{len(segs)}.mp4"
try:
concat_videos(segs, out)
except Exception as e:
raise gr.Error(f"Concat failed (ffmpeg): {e}") from e
# Always refresh last_frame from the new concatenated video so later
# “Last frame from video” extends stay correct.
last = extract_last_frame(out, sess / "last_frame")
dur = video_duration_seconds(out)
state["segments"] = [str(p) for p in segs]
state["video"] = str(out)
state["last_frame"] = str(last)
state["prompt"] = use_prompt
progress(1.0, desc="Extended")
return (
str(out),
str(out),
status_line(state["segments"], dur),
str(last),
state,
)
def do_auto_extend(
target_seconds,
prompt,
extend_prompt,
extend_start,
custom_extend_image,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state: dict,
progress=gr.Progress(track_tqdm=False),
):
if not state or not state.get("video"):
raise gr.Error("Generate a first clip, then Auto-extend.")
target = float(target_seconds)
max_segs = 6
video, download, status, last_preview, state = (
state.get("video"),
state.get("video"),
status_line(state.get("segments", []), video_duration_seconds(state["video"])),
state.get("last_frame"),
state,
)
first = True
while len(state.get("segments", [])) < max_segs:
dur = video_duration_seconds(state["video"])
if dur >= target:
break
progress(
len(state["segments"]) / max_segs,
desc=f"Auto-extend: {dur:.1f}s / {target:.0f}s…",
)
# Custom image only applies to the first auto-extend segment;
# later segments chain from the updated video last frame.
this_start = extend_start if first else EXTEND_START_LAST_FRAME
this_custom = custom_extend_image if first else None
video, download, status, last_preview, state = do_extend(
prompt,
extend_prompt,
this_start,
this_custom,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state,
progress,
)
first = False
time.sleep(0.3)
return video, download, status, last_preview, state
def do_reset(state: dict):
if state and state.get("dir"):
try:
shutil.rmtree(state["dir"], ignore_errors=True)
except Exception:
pass
return None, None, "**Segments:** 0 · **Duration ≈ 0s**", None, {}
def update_estimate(seg_duration, num_planned):
total = estimate_total(seg_duration, num_planned)
return f"Planned length ≈ **{total:.1f}s** ({num_planned} × {seg_duration:.1f}s segments). Cap: 6 segments."
def toggle_custom_extend_image(mode: str):
visible = (mode or "").strip() == EXTEND_START_UPLOAD
return gr.update(visible=visible)
CSS = """
.big-btn button { font-size: 1.15rem !important; min-height: 3rem !important; }
#status-md { font-size: 1rem; }
footer { display: none !important; }
"""
def build_ui():
with gr.Blocks(title="Wan 2.2 Extend — I2V", css=CSS, theme=gr.themes.Soft()) as demo:
gr.Markdown(
"""
# Wan 2.2 Image → Video + Extend
Upload a still, generate a short Wan 2.2 clip on **free ZeroGPU**, then **Extend** from the last frame (or a custom start image) to chain ~10–20s videos.
ZeroGPU is free but **queued and quota-limited** — each segment often takes 1–3 minutes when busy.
"""
)
state = gr.State({})
with gr.Row():
with gr.Column(scale=1):
image = gr.Image(label="Still image", type="pil", height=320)
prompt = gr.Textbox(
label="Generate prompt",
value=DEFAULT_PROMPT,
lines=3,
placeholder="Describe the motion you want…",
)
with gr.Accordion("Extend options", open=True):
extend_prompt = gr.Textbox(
label="Extend prompt",
value="",
lines=3,
placeholder="Leave blank to reuse Generate prompt",
info="Optional. Used only for Extend / Auto-extend.",
)
extend_start = gr.Radio(
choices=[EXTEND_START_LAST_FRAME, EXTEND_START_UPLOAD],
value=EXTEND_START_LAST_FRAME,
label="Extend start",
info="Default keeps current behavior (auto last frame).",
)
custom_extend_image = gr.Image(
label="Custom extend start image (replaces auto last frame)",
type="pil",
height=220,
visible=False,
)
with gr.Row():
seg_duration = gr.Slider(
2.0,
5.0,
value=3.5,
step=0.5,
label="Segment length (seconds)",
info="Each Wan call; Extend adds another segment.",
)
num_planned = gr.Slider(
1,
6,
value=4,
step=1,
label="Target segments",
info="For estimate only (3–5 ≈ 10–20s).",
)
estimate = gr.Markdown(update_estimate(3.5, 4))
target_seconds = gr.Slider(
8,
22,
value=14,
step=1,
label="Auto-extend target (seconds)",
)
with gr.Accordion("Advanced", open=False):
steps = gr.Slider(1, 12, value=4, step=1, label="Inference steps")
quality = gr.Slider(1, 10, value=6, step=1, label="Video quality")
fps = gr.Dropdown(
choices=[16, 32, 64],
value=16,
label="Fluidity FPS",
)
seed = gr.Number(value=42, label="Seed", precision=0)
randomize = gr.Checkbox(value=True, label="Randomize seed")
safe_mode = gr.Checkbox(
value=True,
label="Upstream Safe Mode",
info="Extra ZeroGPU headroom on busy servers.",
)
negative = gr.Textbox(
label="Negative prompt",
value=DEFAULT_NEGATIVE,
lines=2,
)
with gr.Row(elem_classes=["big-btn"]):
btn_gen = gr.Button("Generate", variant="primary", size="lg")
btn_ext = gr.Button("Extend", variant="secondary", size="lg")
with gr.Row(elem_classes=["big-btn"]):
btn_auto = gr.Button("Auto-extend to target", size="lg")
btn_reset = gr.Button("Reset", size="lg")
with gr.Column(scale=1):
video = gr.Video(label="Preview", height=360, autoplay=True)
status = gr.Markdown("**Segments:** 0 · **Duration ≈ 0s**", elem_id="status-md")
last_preview = gr.Image(label="Last frame (next Extend start)", height=180)
download = gr.File(label="Download MP4")
seg_duration.change(update_estimate, [seg_duration, num_planned], estimate)
num_planned.change(update_estimate, [seg_duration, num_planned], estimate)
extend_start.change(toggle_custom_extend_image, [extend_start], [custom_extend_image])
gen_inputs = [
image,
prompt,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state,
]
ext_inputs = [
prompt,
extend_prompt,
extend_start,
custom_extend_image,
seg_duration,
steps,
negative,
seed,
randomize,
quality,
fps,
safe_mode,
state,
]
outs = [video, download, status, last_preview, state]
btn_gen.click(do_generate, gen_inputs, outs)
btn_ext.click(do_extend, ext_inputs, outs)
btn_auto.click(
do_auto_extend,
[target_seconds] + ext_inputs,
outs,
)
btn_reset.click(do_reset, [state], outs)
gr.Markdown(
f"""
### How Extend works
1. **Generate** — still image + Generate prompt → short Wan 2.2 I2V clip (~{3.5}s default).
2. **Extend** — start from the **last frame** (or a **custom uploaded image**), optional separate **Extend prompt**, run I2V again, **concatenate**.
3. Repeat until you hit ~10–20s (or use **Auto-extend**). Soft cap: **6 segments**.
### Limits
- Powered by public ZeroGPU Space `{UPSTREAM_SPACE}` — free, but **queued / quota-limited**, not unlimited.
- Single segments are typically a few seconds; longer clips = chained Extends.
- Hugging Face Terms of Service apply to hosted Spaces.
"""
)
return demo
if __name__ == "__main__":
demo = build_ui()
demo.queue(default_concurrency_limit=1).launch(server_name="0.0.0.0", server_port=7860)