multimodalart's picture
multimodalart HF Staff
Update app.py
22c83c7 verified
Raw History Blame Contribute Delete
18.9 kB
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402 (must precede torch / CUDA imports)
import torch # noqa: E402
import numpy as np # noqa: E402
import tempfile # noqa: E402
import subprocess # noqa: E402
import gradio as gr # noqa: E402
from PIL import Image # noqa: E402
from huggingface_hub import snapshot_download # noqa: E402
from safetensors.torch import load_file # noqa: E402
from diffusers import WanAnimate2Pipeline, WanAnimate2Transformer3DModel # noqa: E402
from diffusers.utils import export_to_video, load_image # noqa: E402
# ---------------------------------------------------------------------------
# ZeroGPU cannot host torch.compile / dynamo inside the forked GPU worker.
# The Wan-Animate-2 transformer builds its in-context BlockMask with
# `flex_attention.create_block_mask(..., _compile=True)`, which triggers a
# dynamo compile and aborts on ZeroGPU. Force the eager (`_compile=False`)
# path so the mask is built without compilation; flex_attention itself is
# then called eagerly, which is supported.
# ---------------------------------------------------------------------------
import torch.nn.attention.flex_attention as _flex_mod # noqa: E402
_orig_create_block_mask = _flex_mod.create_block_mask
def _create_block_mask_no_compile(*args, **kwargs):
kwargs["_compile"] = False
return _orig_create_block_mask(*args, **kwargs)
_flex_mod.create_block_mask = _create_block_mask_no_compile
# Also patch the reference already imported into the transformer module.
try:
import diffusers.models.transformers.transformer_wan_animate_2 as _wa2 # noqa: E402
_wa2.create_block_mask = _create_block_mask_no_compile
except Exception as _e: # pragma: no cover
print(f"[patch] could not patch transformer_wan_animate_2.create_block_mask: {_e!r}")
# The Wan-Animate-2 in-context self-attention runs on the FLEX backend with a
# custom `BlockMask`. On ZeroGPU we cannot `torch.compile` flex_attention (the
# forked GPU worker has no dynamo daemon, and this torch build's dynamo also
# chokes tracing `query.device != key.device`), and the eager flex fallback
# decomposes into a dense `Q@K^T` (`math_attention`) that OOMs the allocator on
# a 14B video DiT. Instead we replace the FLEX backend with a memory-efficient
# SDPA path: materialise the BlockMask's boolean mask ONCE (Q*KV bytes, cheap)
# and hand it to `scaled_dot_product_attention`, whose mem-efficient kernel
# never builds the full scores matrix. Same masking semantics, no compile.
from torch.nn.attention.flex_attention import BlockMask as _BlockMask # noqa: E402
import torch.nn.functional as _F # noqa: E402
import diffusers.models.attention_dispatch as _attn_dispatch # noqa: E402
_dense_mask_cache = {}
def _blockmask_to_dense_bool(block_mask, seq_len_q, seq_len_kv, device):
key = (id(block_mask), seq_len_q, seq_len_kv)
cached = _dense_mask_cache.get(key)
if cached is not None:
return cached
mask_mod = block_mask.mask_mod
zero = torch.zeros((), dtype=torch.long, device=device)
kv_idx = torch.arange(seq_len_kv, device=device)
# Build the [Q, KV] boolean mask row-block by row-block so we never hold
# multiple Q*KV int64 temporaries at once (the mask_mod uses torch.where on
# int64 index grids, which spikes memory badly at high resolution).
dense = torch.empty((seq_len_q, seq_len_kv), dtype=torch.bool, device=device)
q_chunk = max(1, min(seq_len_q, 512))
for q0 in range(0, seq_len_q, q_chunk):
q1 = min(seq_len_q, q0 + q_chunk)
rows = q1 - q0
qg = torch.arange(q0, q1, device=device).view(rows, 1).expand(rows, seq_len_kv)
kg = kv_idx.view(1, seq_len_kv).expand(rows, seq_len_kv)
dense[q0:q1] = mask_mod(zero, zero, qg, kg).to(torch.bool)
dense = dense.view(1, 1, seq_len_q, seq_len_kv)
if len(_dense_mask_cache) > 4:
_dense_mask_cache.clear()
_dense_mask_cache[key] = dense
return dense
def _flex_backend_sdpa(
query, key, value, attn_mask=None, is_causal=False, scale=None,
enable_gqa=False, return_lse=False, _parallel_config=None,
):
# query/key/value: [B, seq, H, D] -> SDPA wants [B, H, seq, D]
batch_size, seq_len_q, _, _ = query.shape
seq_len_kv = key.shape[1]
q = query.permute(0, 2, 1, 3)
k = key.permute(0, 2, 1, 3)
v = value.permute(0, 2, 1, 3)
sdpa_mask = None
if isinstance(attn_mask, _BlockMask):
sdpa_mask = _blockmask_to_dense_bool(attn_mask, seq_len_q, seq_len_kv, query.device)
elif attn_mask is not None and torch.is_tensor(attn_mask):
sdpa_mask = attn_mask
out = _F.scaled_dot_product_attention(
q, k, v, attn_mask=sdpa_mask, is_causal=is_causal, scale=scale, enable_gqa=enable_gqa,
)
out = out.permute(0, 2, 1, 3)
if return_lse:
return out, None
return out
# Re-register the FLEX backend to point at the SDPA-based implementation.
_attn_dispatch._native_flex_attention = _flex_backend_sdpa
try:
_reg = _attn_dispatch._AttentionBackendRegistry
_reg._backends[_attn_dispatch.AttentionBackendName.FLEX] = _flex_backend_sdpa
print("[patch] FLEX backend replaced with memory-efficient SDPA path")
except Exception as _e: # pragma: no cover
print(f"[patch] could not re-register FLEX backend: {_e!r}")
MODEL_ID = "Wan-AI/Wan2.2-Animate-2-14B-Distilled-Diffusers"
DEFAULT_NEGATIVE_PROMPT = (
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,"
"低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,"
"毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
)
# ---------------------------------------------------------------------------
# The published checkpoint uses the ORIGINAL Wan-Animate-2 weight naming
# (from the author diffusers PR #14412 / the reference implementation), while
# the refactored diffusers branch (#14413) we build against renames the block
# internals to diffusers conventions. Remap the state dict on load so the
# ZeroGPU-friendly refactored model (SDPA / flex attention, no hard flash_attn
# and no torch.compile at import) can consume the weights.
# ---------------------------------------------------------------------------
def _remap_transformer_key(key: str) -> str:
if not key.startswith("blocks."):
return key
parts = key.split(".")
# blocks.N.block.<rest...> -> blocks.N.<mapped-rest>
if len(parts) < 4 or parts[2] != "block":
return key
n = parts[1]
rest = parts[3:]
sub = rest[0] # self_attn | cross_attn | ffn | norm3 | modulation
if sub in ("self_attn", "cross_attn"):
proj = rest[1]
tail = ".".join(rest[2:])
proj_map = {
"q": "to_q",
"k": "to_k",
"v": "to_v",
"o": "to_out.0",
"k_img": "add_k_proj",
"v_img": "add_v_proj",
"norm_q": "norm_q",
"norm_k": "norm_k",
"norm_k_img": "norm_added_k",
}
mapped = proj_map.get(proj, proj)
new = f"blocks.{n}.{sub}.{mapped}"
return f"{new}.{tail}" if tail else new
# ffn.*, norm3.*, modulation -> passthrough after dropping `block`
return f"blocks.{n}." + ".".join(rest)
def _load_transformer():
ckpt_dir = snapshot_download(MODEL_ID, allow_patterns=["transformer/*"])
tdir = os.path.join(ckpt_dir, "transformer")
# Load + merge all shards.
import glob
remapped = {}
for shard in sorted(glob.glob(os.path.join(tdir, "*.safetensors"))):
sd = load_file(shard)
for k, v in sd.items():
remapped[_remap_transformer_key(k)] = v.to(torch.bfloat16)
del sd
model = WanAnimate2Transformer3DModel.from_config(
WanAnimate2Transformer3DModel.load_config(tdir)
).to(torch.bfloat16)
missing, unexpected = model.load_state_dict(remapped, strict=False)
missing = [m for m in missing if "kv_cache" not in m]
if missing:
raise RuntimeError(f"Missing transformer keys after remap: {missing[:20]} (total {len(missing)})")
if unexpected:
raise RuntimeError(f"Unexpected transformer keys after remap: {unexpected[:20]} (total {len(unexpected)})")
return model
# ---------------------------------------------------------------------------
# Load the full pipeline at module scope, eagerly to CUDA (ZeroGPU packs it).
# ---------------------------------------------------------------------------
transformer = _load_transformer()
pipe = WanAnimate2Pipeline.from_pretrained(
MODEL_ID, transformer=transformer, torch_dtype=torch.bfloat16
)
pipe.to("cuda")
# ---------------------------------------------------------------------------
# AoTI (Ahead-of-Time Inductor) acceleration.
#
# The precompiled artifacts are produced offline by a dedicated compile Space
# (multimodalart/wan2-2-animate-2-aoti-compile) and published to the model repo
# below, so this serving Space pays no inline compile cost at startup.
#
# We compile the repeated block's feed-forward sub-graph (`block.ffn`) rather
# than the whole WanAnimate2TransformerBlock: the block's self-attention calls
# `rope_apply`, which reads `grid_sizes.tolist()` to drive tensor reshapes — a
# data-dependent op that `torch.export` refuses. The FFN (Linear->GELU->Linear)
# is a clean tensor->tensor graph, is a meaningful slice of per-block compute,
# and is architecturally identical across all 40 blocks, so one compiled graph
# (with per-block weights streamed in as runtime inputs) patches every block.
# Falls back to eager on any error so the demo keeps running regardless.
# ---------------------------------------------------------------------------
AOTI_REPO = "multimodalart/Wan2.2-Animate-2-14B-Distilled-aoti"
AOTI_FFN_DIR = "WanAnimate2TransformerBlockFFN"
try:
from pathlib import Path as _Path
from spaces.zero.torch.aoti import aoti_load_from_module_dir
_aoti_root = _Path(snapshot_download(AOTI_REPO, allow_patterns=[f"{AOTI_FFN_DIR}/*"]))
_ffn_pkg_dir = _aoti_root / AOTI_FFN_DIR
if (_ffn_pkg_dir / "package.pt2").exists():
_ffns = [blk.ffn for blk in pipe.transformer.blocks]
aoti_load_from_module_dir(_ffns, _ffn_pkg_dir)
print(f"[aoti] patched {len(_ffns)} block.ffn modules from {AOTI_REPO}/{AOTI_FFN_DIR}")
else:
print(f"[aoti] {AOTI_FFN_DIR}/package.pt2 not found in {AOTI_REPO}; running eager")
except Exception as _e: # pragma: no cover
print(f"[aoti] AoTI load failed ({_e!r}); running eager")
def _trim_video(src_path: str, max_seconds: float, target_fps: int = 24) -> str:
"""Trim the driving video to the first `max_seconds` at `target_fps`.
The pipeline animates the whole clip in 81-frame segments, so a long clip
is very slow. Trimming keeps a single-request demo tractable. Falls back to
the original path if ffmpeg is unavailable.
"""
out = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name
try:
subprocess.run(
[
"ffmpeg", "-y", "-i", src_path,
"-t", str(max_seconds),
"-r", str(target_fps),
"-an",
"-c:v", "libx264", "-pix_fmt", "yuv420p",
out,
],
check=True,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
return out
except Exception as e:
print(f"[trim] ffmpeg failed ({e!r}); using original video")
return src_path
def _estimate_duration(image, driving_video, prompt, max_seconds=2.0, height=480,
width=384, num_inference_steps=6, *args, **kwargs):
# The FLEX self-attention runs on a memory-efficient (but not compiled) SDPA
# path whose cost grows ~quadratically with the token count (h*w), and
# linearly with steps and segments. Anchored on two live measurements:
# 320x480, 4 steps, 1 seg -> ~23 s
# 640x480, 8 steps, 1 seg -> ~237 s
# Model: per_seg = overhead + k * steps * (pixels / P0)^1.6 .
try:
segments = max(1, int(np.ceil(float(max_seconds) / 3.4)))
except Exception:
segments = 1
try:
pixels = float(height) * float(width)
steps = max(1, int(num_inference_steps))
except Exception:
pixels, steps = 480 * 384, 6
P0 = 320.0 * 480.0
per_step = 3.25 * (pixels / P0) ** 3.1 # seconds per step
per_segment = steps * per_step # dominant per-segment cost
total = 10.0 + segments * per_segment # + model/setup overhead
return int(min(1500, total * 1.2 + 15))
@spaces.GPU(duration=_estimate_duration, size="xlarge")
def animate(
image,
driving_video,
prompt,
max_seconds: float = 2.0,
height: int = 480,
width: int = 384,
num_inference_steps: int = 6,
guidance_scale: float = 1.0,
sample_shift: float = 5.0,
negative_prompt: str = DEFAULT_NEGATIVE_PROMPT,
seed: int = 0,
progress=gr.Progress(track_tqdm=True),
):
"""Animate a reference character image with the motion from a driving video.
Args:
image: Reference character image (a person / character to animate).
driving_video: Driving video whose motion is transferred to the character.
prompt: Text description of the character appearance and background.
max_seconds: How many seconds of the driving video to animate.
height: Output height (multiple of 16).
width: Output width (multiple of 16).
num_inference_steps: Denoising steps per segment.
guidance_scale: Classifier-free guidance scale (1.0 = off, faster).
sample_shift: Flow-matching sigma shift.
negative_prompt: Negative prompt for guidance (used when guidance > 1).
seed: Random seed.
Returns:
Path to the generated animated video (mp4).
"""
if image is None:
raise gr.Error("Please provide a reference character image.")
if driving_video is None:
raise gr.Error("Please provide a driving video.")
if isinstance(image, str):
image = load_image(image)
if not isinstance(image, Image.Image):
image = Image.fromarray(np.asarray(image))
image = image.convert("RGB")
height = int(height) - (int(height) % 16)
width = int(width) - (int(width) % 16)
trimmed = _trim_video(driving_video, float(max_seconds), target_fps=24)
generator = torch.Generator(device="cuda").manual_seed(int(seed))
try:
output = pipe(
image=image,
driving_video=trimmed,
prompt=prompt or "",
negative_prompt=negative_prompt,
height=height,
width=width,
fps=24,
num_inference_steps=int(num_inference_steps),
guidance_scale=float(guidance_scale),
sample_shift=float(sample_shift),
seed=int(seed),
generator=generator,
output_type="np",
)
except Exception as exc: # surface the real traceback to the client/logs
import traceback
tb = traceback.format_exc()
print("[animate] inference failed:\n" + tb, flush=True)
raise gr.Error(f"{type(exc).__name__}: {exc}\n{tb[-1500:]}")
frames = output.frames[0]
out_path = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name
export_to_video(frames, out_path, fps=24)
return out_path
CSS = """
#col-container { max-width: 1200px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# Wan2.2-Animate-2-14B
Animate a **reference character image** with the motion from a
**driving video**. Powered by
[`Wan-AI/Wan2.2-Animate-2-14B-Distilled-Diffusers`](https://huggingface.co/Wan-AI/Wan2.2-Animate-2-14B-Distilled-Diffusers)
via 🧨 diffusers.
"""
)
with gr.Row():
with gr.Column():
image = gr.Image(label="Reference character image", type="filepath", height=340)
driving_video = gr.Video(label="Driving video (motion source)", height=340)
prompt = gr.Textbox(
label="Prompt",
placeholder="Describe the character and background…",
value="a person",
lines=2,
)
run = gr.Button("Animate", variant="primary")
with gr.Column():
output_video = gr.Video(label="Animated result", height=520)
with gr.Accordion("Advanced settings", open=False):
max_seconds = gr.Slider(
label="Seconds of driving video to animate",
minimum=1.0, maximum=5.0, step=0.5, value=2.0,
)
with gr.Row():
height = gr.Slider(label="Height", minimum=320, maximum=640, step=16, value=480)
width = gr.Slider(label="Width", minimum=320, maximum=640, step=16, value=384)
with gr.Row():
num_inference_steps = gr.Slider(
label="Inference steps", minimum=4, maximum=20, step=1, value=6,
)
guidance_scale = gr.Slider(
label="Guidance scale (1 = off)", minimum=1.0, maximum=6.0, step=0.5, value=1.0,
)
sample_shift = gr.Slider(
label="Sample shift", minimum=1.0, maximum=12.0, step=0.5, value=5.0,
)
negative_prompt = gr.Textbox(
label="Negative prompt", value=DEFAULT_NEGATIVE_PROMPT, lines=2,
)
seed = gr.Number(label="Seed", value=0, precision=0)
gr.Examples(
examples=[
["examples/animate_character.jpeg", "examples/animate_driving.mp4", "a person dancing"],
["examples/replace_character.jpeg", "examples/replace_driving.mp4", "a person"],
],
inputs=[image, driving_video, prompt],
outputs=output_video,
fn=animate,
cache_examples=True,
cache_mode="lazy",
)
run.click(
fn=animate,
inputs=[
image, driving_video, prompt, max_seconds, height, width,
num_inference_steps, guidance_scale, sample_shift, negative_prompt, seed,
],
outputs=output_video,
api_name="animate",
)
if __name__ == "__main__":
demo.launch(mcp_server=True)