Spaces:
Running on Zero
Running on Zero
Download app.py from hugging-apps/wan2-2-animate-2-14b: direct link, hf CLI and curl.
- Browser
- Download file 18.9 kB
-
https://huggingface.co/spaces/hugging-apps/wan2-2-animate-2-14b/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/wan2-2-animate-2-14b/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/wan2-2-animate-2-14b/resolve/main/app.py
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)) | |
| 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) | |