Wan2.2-Animate2 / app.py
Transfer Bot
Moved to Hugging Face automatically
a8d7e70
Raw History Blame Contribute Delete
96.9 kB
# ==========================================================================
# Wan2.2-Animate-2-14B (Wan-Animate-2) — اسپیس ZeroGPU با رندر چندمرحله‌ای قابل ادامه
# + حالت «جایگزینی شخصیت» (Mix) با Wan2.2-Animate-14B — بخش mix_engine.py
# API بدون تغییر: animate_scene(input_video, edited_frame, resolution_choice, prompt_text, state_file)
# resolution_choice: "720p" | "1080p" (انتقال حرکت) یا "mix-480p" | "mix-720p" (جایگزینی شخصیت)
# ==========================================================================
# علت خطای «RuntimeError: No CUDA GPUs are available» و راه‌حل:
# مدل‌های T5 و CLIP روی CPU داخل پروسه‌ی اصلی اجرا می‌شدند. transformers حین اجرا
# torch.cuda.is_current_stream_capturing() را صدا می‌زند و همین، وضعیت CUDA پروسه‌ی
# اصلی را (که هیچ GPU واقعی ندارد) خراب می‌کرد؛ در نتیجه هر پنجره‌ی GPU که ZeroGPU از
# آن fork می‌کرد، هنگام راه‌اندازی CUDA شکست می‌خورد (هر دو کیفیت).
# حالا همه‌ی محاسبات CPU در پروسه‌ی مستقل cpu_encoder.py انجام می‌شود و پروسه‌ی اصلی
# هیچ کد مدلی اجرا نمی‌کند؛ علاوه بر این، آن فراخوانی در پروسه‌ی اصلی بی‌اثر شده است.
#
# سرعت: مدل Distilled (۱۰ گام، بدون CFG ≈ ۸ برابر سریع‌تر)، کارت کامل (xlarge) و پنجره‌های
# GPU بلند (تا ۲۰ دقیقه) که مدتشان از روی کار باقی‌مانده تخمین زده می‌شود؛ ویدیوهای معمولی
# در یک مرحله کامل می‌شوند.
# ==========================================================================
import os
import gc
import re
import sys
import copy
import json
import math
import time
import uuid
import random
import select
import contextvars
import shutil
import hashlib
import tempfile
import threading
import traceback
import subprocess
from concurrent.futures import ThreadPoolExecutor, wait as futures_wait
import numpy as np
import cv2
import spaces
import torch
import gradio as gr
from huggingface_hub import snapshot_download
# لایه‌ی دوم محافظت: هیچ کدی در پروسه‌ی اصلی نباید وضعیت CUDA را لمس کند.
# (CUDA Graph در این برنامه استفاده نمی‌شود، پس False همیشه پاسخ درست است.)
def _not_capturing():
return False
torch.cuda.is_current_stream_capturing = _not_capturing
try:
import torch.cuda.graphs as _torch_cuda_graphs
_torch_cuda_graphs.is_current_stream_capturing = _not_capturing
except Exception:
pass
import diffusers
from diffusers import AutoencoderKLWan, WanAnimate2Transformer3DModel
from diffusers.models.transformers.transformer_wan_animate_2 import WanAnimate2KVCache
from diffusers.modular_pipelines.wan_animate_2.encoders import encode_vae, get_i2v_mask
from diffusers.modular_pipelines.wan_animate_2.denoise import decode_vae
from diffusers.utils.torch_utils import randn_tensor
try:
import transformers.utils.import_utils as _tf_import_utils
_tf_import_utils.is_cuda_stream_capturing = _not_capturing
except Exception:
pass
import animate2_engine as eng
import mix_engine as mix
print(sys.version)
print("torch", torch.__version__, "| diffusers", diffusers.__version__)
APP_DIR = os.path.dirname(os.path.abspath(__file__))
ENCODER_SCRIPT = os.path.join(APP_DIR, "cpu_encoder.py")
# --------------------------------------------------------------------------
# تنظیمات اصلی
# --------------------------------------------------------------------------
MODEL_REPO = "Wan-AI/Wan2.2-Animate-2-14B-Distilled-Diffusers"
IS_DISTILLED = "Distilled" in MODEL_REPO
NUM_INFERENCE_STEPS = 10 if IS_DISTILLED else 40
GUIDANCE_SCALE = 1.0 if IS_DISTILLED else 3.0
CFG_PASSES = 2 if GUIDANCE_SCALE > 1.0 else 1
RESOLUTIONS = ["480p", "720p"]
# 1080p فقط برای ادامه‌ی فایل‌های وضعیت قدیمی نگه داشته شده؛ درخواست جدید 1080p به 720p تبدیل می‌شود.
RESOLUTION_AREA = {"480p": 832 * 480, "720p": 1280 * 720, "1080p": 1920 * 1080}
# مدت هر پنجره‌ی GPU: از روی کار باقی‌مانده تخمین زده می‌شود و حداکثر ۲۰ دقیقه است.
# (با xlarge هر ثانیه دو برابر از سهمیه کم می‌شود و سقف هر درخواست برای حساب PRO،
# ۴۰ دقیقه سهمیه است؛ پس ۲۰ دقیقه بیشترین مقدار مجاز برای هر مرحله است.)
GPU_SIZE = "xlarge" # کارت گرافیک کامل (۹۶ گیگابایت)
MAX_WINDOW_SECONDS = 1200
MIN_WINDOW_SECONDS = {"480p": 90, "720p": 120, "1080p": 240}
WINDOW_ESTIMATE_FACTOR = 1.35
WINDOW_INIT_ALLOWANCE = 45 # انتقال وزن‌ها به GPU در شروع هر پنجره
GPU_SAFETY_SECONDS = 20
MAX_PIXEL_DISK_BYTES = 6 * 1024 ** 3
DISK_MARGIN_BYTES = 4 * 1024 ** 3
PROGRESS_POLL_SECONDS = 15
CHECKPOINT_EVERY_SECONDS = 60
MAX_OOM_RETRIES = 4
OOM_EXTRA_GB = 8.0
# برآورد اولیه‌ی زمان هر مرحله (ثانیه)؛ از پنجره‌ی دوم به بعد با زمان واقعی جایگزین می‌شود
DEFAULT_TIMINGS = {
"480p": {"encode": 5.0, "extract": 3.5, "fwd": 2.0, "decode": 6.0},
"720p": {"encode": 10.0, "extract": 8.0, "fwd": 5.0, "decode": 12.0},
"1080p": {"encode": 25.0, "extract": 25.0, "fwd": 22.0, "decode": 35.0},
}
WINDOW_OVERHEAD_SECONDS = {"480p": 5.0, "720p": 8.0, "1080p": 15.0} # انکود تصویر مرجع در شروع هر پنجره
STATE_FORMAT = "wan-animate-2/v1"
MAX_VIDEO_LONG_SIDE = 1920
ENCODER_STARTUP_TIMEOUT = 1800
ENCODER_PREPARE_TIMEOUT = 900
ENCODER_SEGMENTS_TIMEOUT = 600
ENCODER_SEGMENT_EXTRA_TIMEOUT = 180
OUTPUT_DIR = os.path.join(tempfile.gettempdir(), "wan_anim2_outputs")
OUTPUT_MAX_AGE_SECONDS = 3 * 3600
WORKDIR_PREFIX = "wan_anim2_"
os.makedirs(OUTPUT_DIR, exist_ok=True)
# --------------------------------------------------------------------------
# بارگذاری مدل جدید (مدل قدیمی Wan2.2-Animate-14B دیگر استفاده نمی‌شود)
# --------------------------------------------------------------------------
model_dir = snapshot_download(repo_id=MODEL_REPO)
# --------------------------------------------------------------------------
# دانلود فایل‌های حالت «جایگزینی شخصیت» (Wan2.2-Animate-14B) در پس‌زمینه
# (همزمان با بارگذاری مدل اصلی؛ اگر ناموفق باشد حالت اصلی بدون مشکل کار می‌کند)
# --------------------------------------------------------------------------
MIX_ASSETS = {"model_dir": None, "aux_dir": None, "scheduler": None, "error": None}
_mix_assets_lock = threading.Lock()
_mix_assets_ready = threading.Event()
def _download_mix_assets():
with _mix_assets_lock:
if MIX_ASSETS["scheduler"] is not None:
return True
try:
print("[mix] downloading replacement-mode assets...", flush=True)
mdir = snapshot_download(repo_id=mix.MIX_MODEL_REPO, allow_patterns=mix.MIX_MODEL_PATTERNS)
adir = snapshot_download(repo_id=mix.MIX_AUX_REPO, allow_patterns=mix.MIX_AUX_PATTERNS)
with open(os.path.join(mdir, "scheduler", "scheduler_config.json"), "r", encoding="utf-8") as sf:
sched_cls = getattr(diffusers, json.load(sf)["_class_name"])
sched = sched_cls.from_pretrained(mdir, subfolder="scheduler")
MIX_ASSETS.update(model_dir=mdir, aux_dir=adir, scheduler=sched, error=None)
print("[mix] replacement-mode assets ready", flush=True)
return True
except Exception as e:
MIX_ASSETS["error"] = f"{type(e).__name__}: {e}"
print(f"[mix] asset download failed: {MIX_ASSETS['error']}", flush=True)
return False
finally:
_mix_assets_ready.set()
threading.Thread(target=_download_mix_assets, daemon=True).start()
# --------------------------------------------------------------------------
# کلاینت پروسه‌ی انکودر CPU
# --------------------------------------------------------------------------
class EncoderTaskError(RuntimeError):
"""خطای منطقی در پردازش ورودی (مثلاً تصویر خراب) — تکرار درخواست کمکی نمی‌کند."""
class CpuEncoder:
def __init__(self, model_path):
self.model_path = model_path
self.lock = threading.Lock()
self.proc = None
self.fixed = None
self.fixed_path = os.path.join(tempfile.gettempdir(), f"wan_fixed_prompts_{os.getpid()}.pt")
def _kill(self):
proc, self.proc = self.proc, None
self.fixed = None
if proc is None:
return
try:
proc.kill()
proc.wait(timeout=30)
except Exception:
pass
def _readline(self, timeout):
deadline = time.time() + timeout
while True:
if self.proc.poll() is not None:
raise RuntimeError(f"پروسه‌ی انکودر بسته شد (کد {self.proc.returncode})")
wait = min(5.0, deadline - time.time())
if wait <= 0:
raise TimeoutError("پاسخ پروسه‌ی انکودر در زمان مقرر نرسید")
ready, _, _ = select.select([self.proc.stdout], [], [], wait)
if ready:
line = self.proc.stdout.readline()
if not line:
continue
return json.loads(line)
def _ensure(self):
if self.proc is not None and self.proc.poll() is None and self.fixed is not None:
return
self._kill()
env = dict(os.environ)
env["CUDA_VISIBLE_DEVICES"] = ""
env["PYTHONUNBUFFERED"] = "1"
print("[encoder] starting CPU encoder process...", flush=True)
self.proc = subprocess.Popen(
[sys.executable, ENCODER_SCRIPT, self.model_path, self.fixed_path],
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=None,
text=True,
bufsize=1,
encoding="utf-8",
env=env,
cwd=APP_DIR,
)
msg = self._readline(ENCODER_STARTUP_TIMEOUT)
if not msg.get("ready"):
raise RuntimeError(f"پروسه‌ی انکودر آماده نشد: {msg}")
self.fixed = torch.load(self.fixed_path, map_location="cpu")
print("[encoder] CPU encoder ready", flush=True)
def warmup(self):
with self.lock:
try:
self._ensure()
except Exception as e:
print(f"[encoder] startup failed (will retry on first request): {e}", flush=True)
self._kill()
def fixed_embeds(self):
with self.lock:
for attempt in range(2):
try:
self._ensure()
return self.fixed
except Exception as e:
print(f"[encoder] start attempt {attempt + 1} failed: {e}", flush=True)
self._kill()
raise gr.Error("راه‌اندازی انکودر متن/تصویر ناموفق بود؛ چند لحظه بعد دوباره تلاش کنید.")
def request(self, payload, timeout):
with self.lock:
last_error = None
for attempt in range(2):
try:
self._ensure()
self.proc.stdin.write(json.dumps(payload, ensure_ascii=False) + "\n")
self.proc.stdin.flush()
msg = self._readline(timeout)
if not msg.get("ok"):
raise EncoderTaskError(msg.get("error") or "خطای نامشخص")
return msg
except EncoderTaskError:
raise
except Exception as e:
last_error = e
print(f"[encoder] request '{payload.get('cmd')}' attempt {attempt + 1} failed: {e}", flush=True)
self._kill()
raise RuntimeError(f"پروسه‌ی انکودر پاسخ نداد: {last_error}")
ENCODER = CpuEncoder(model_dir)
_encoder_warmup = threading.Thread(target=ENCODER.warmup, daemon=True)
_encoder_warmup.start()
with open(os.path.join(model_dir, "scheduler", "scheduler_config.json"), "r", encoding="utf-8") as f:
_sched_cls = getattr(diffusers, json.load(f)["_class_name"])
base_scheduler = _sched_cls.from_pretrained(model_dir, subfolder="scheduler")
transformer = WanAnimate2Transformer3DModel.from_pretrained(model_dir, subfolder="transformer", dtype=torch.bfloat16)
transformer.eval()
kv_policy = eng.install_exact_attention(transformer)
transformer.to("cuda")
vae = AutoencoderKLWan.from_pretrained(model_dir, subfolder="vae", dtype=torch.bfloat16)
vae.eval()
vae.to("cuda")
# انکودرهای متن و تصویر در پروسه‌ی مستقل روی CPU می‌مانند تا ثانیه‌های GPU فقط صرف رندر شود
_encoder_warmup.join(timeout=ENCODER_STARTUP_TIMEOUT)
print("Model ready:", MODEL_REPO, flush=True)
# --------------------------------------------------------------------------
# ابزارهای ویدیو/صدا/فایل (بدون هیچ محاسبه‌ی torch)
# --------------------------------------------------------------------------
def normalize_resolution(value):
"""انتقال حرکت: 480p یا 720p (مقدار قدیمی 1080p هم به 720p تبدیل می‌شود)."""
v = str(value or "").lower()
if "480" in v:
return "480p"
return "720p"
def run_ffmpeg(cmd):
try:
subprocess.run(cmd, check=True, capture_output=True, text=True)
except subprocess.CalledProcessError as e:
raise RuntimeError(f"ffmpeg failed ({e.returncode}): {e.stderr.strip()[-500:]}")
def transcode_to_model_fps(input_path, output_path, fps=None):
fps = fps or eng.OUTPUT_FPS
scale = (
f"scale=trunc(iw*min(1\\,{MAX_VIDEO_LONG_SIDE}/max(iw\\,ih))/2)*2:"
f"trunc(ih*min(1\\,{MAX_VIDEO_LONG_SIDE}/max(iw\\,ih))/2)*2"
)
run_ffmpeg([
"ffmpeg", "-nostdin", "-hide_banner", "-y", "-i", input_path,
"-an", "-vf", f"fps={fps},{scale}",
"-c:v", "libx264", "-pix_fmt", "yuv420p", "-preset", "veryfast", "-crf", "16",
output_path,
])
def extract_audio(video_path, output_wav_path):
cmd = ["ffmpeg", "-nostdin", "-y", "-loglevel", "error", "-i", video_path, "-vn",
"-acodec", "pcm_s16le", "-ac", "2", output_wav_path]
try:
subprocess.run(cmd, check=True, capture_output=True, text=True)
return os.path.exists(output_wav_path) and os.path.getsize(output_wav_path) > 1024
except subprocess.CalledProcessError:
return False
def count_frames(video_path):
cap = cv2.VideoCapture(video_path)
n = 0
while cap.grab():
n += 1
cap.release()
return n
def encode_frames_to_mp4(frames_uint8, output_path, fps=None):
fps = fps or eng.OUTPUT_FPS
t, h, w, _ = frames_uint8.shape
h2, w2 = h - (h % 2), w - (w % 2)
frames_uint8 = np.ascontiguousarray(frames_uint8[:, :h2, :w2])
proc = subprocess.Popen(
["ffmpeg", "-nostdin", "-hide_banner", "-loglevel", "error", "-y",
"-f", "rawvideo", "-pix_fmt", "rgb24", "-s", f"{w2}x{h2}", "-r", str(fps), "-i", "-",
"-c:v", "libx264", "-pix_fmt", "yuv420p", "-preset", "medium", "-crf", "16", output_path],
stdin=subprocess.PIPE, stderr=subprocess.PIPE,
)
_, err = proc.communicate(frames_uint8.tobytes())
if proc.returncode != 0:
raise RuntimeError(f"ffmpeg encode failed: {err.decode(errors='ignore')[-500:]}")
def concat_mp4s(paths, output_path):
list_file = tempfile.NamedTemporaryFile(suffix=".txt", delete=False, mode="w")
for p in paths:
list_file.write(f"file '{os.path.abspath(p)}'\n")
list_file.close()
try:
run_ffmpeg(["ffmpeg", "-nostdin", "-y", "-f", "concat", "-safe", "0", "-i", list_file.name,
"-c", "copy", "-movflags", "+faststart", output_path])
finally:
os.remove(list_file.name)
def combine_video_and_audio(video_path, audio_path, output_path):
run_ffmpeg(["ffmpeg", "-nostdin", "-y", "-loglevel", "error", "-i", video_path, "-i", audio_path,
"-map", "0:v:0", "-map", "1:a:0", "-c:v", "copy", "-c:a", "aac", "-shortest",
"-movflags", "+faststart", output_path])
def write_bytes(data, suffix, workdir):
path = os.path.join(workdir, f"{uuid.uuid4().hex}{suffix}")
with open(path, "wb") as f:
f.write(data)
return path
def _npy_to_mp4(npy_path, mp4_path, fps=None):
"""فریم‌های یک بخش را به mp4 تبدیل می‌کند (فایل npy فقط بعد از موفقیت پاک می‌شود)."""
frames = np.load(npy_path)
tmp = mp4_path + ".tmp.mp4"
encode_frames_to_mp4(frames, tmp, fps)
os.replace(tmp, mp4_path)
try:
os.remove(npy_path)
except OSError:
pass
def atomic_torch_save(obj, path):
tmp = f"{path}.{uuid.uuid4().hex[:6]}.tmp"
torch.save(obj, tmp)
os.replace(tmp, path)
def prune_old_files():
"""جلوگیری از پر شدن دیسک اسپیس با فایل‌های موقت قدیمی."""
now = time.time()
try:
for name in os.listdir(OUTPUT_DIR):
path = os.path.join(OUTPUT_DIR, name)
try:
if now - os.path.getmtime(path) > OUTPUT_MAX_AGE_SECONDS:
os.remove(path)
except OSError:
pass
tmp_root = tempfile.gettempdir()
for name in os.listdir(tmp_root):
if not name.startswith(WORKDIR_PREFIX):
continue
path = os.path.join(tmp_root, name)
try:
if os.path.isdir(path) and now - os.path.getmtime(path) > 6 * 3600:
shutil.rmtree(path, ignore_errors=True)
except OSError:
pass
except Exception as e:
print(f"[cleanup] {e}", flush=True)
# --------------------------------------------------------------------------
# ساخت وضعیت اولیه‌ی یک کار جدید (پیش‌پردازش در پروسه‌ی انکودر CPU)
# --------------------------------------------------------------------------
def init_state(input_video, edited_frame, resolution, prompt_text, workdir):
if not input_video:
raise gr.Error("لطفاً ویدیوی مرجع (حرکت) را آپلود کنید.")
if not edited_frame:
raise gr.Error("لطفاً تصویر مرجع را آپلود کنید.")
processed_video = os.path.join(workdir, "driving_24fps.mp4")
try:
transcode_to_model_fps(input_video, processed_video)
except Exception as e:
print(f"[job] transcode failed: {e}", flush=True)
raise gr.Error("ویدیوی ورودی قابل خواندن نیست؛ لطفاً فایل ویدیوی دیگری امتحان کنید.")
real_frame_len = count_frames(processed_video)
if real_frame_len < 1:
raise gr.Error("ویدیوی ورودی هیچ فریم قابل خواندنی ندارد.")
audio_path = os.path.join(workdir, "audio.wav")
audio_bytes = None
if extract_audio(input_video, audio_path):
with open(audio_path, "rb") as f:
audio_bytes = f.read()
prompt_text = (prompt_text or "").strip()
prep_path = os.path.join(workdir, "prepared.pt")
try:
prep = ENCODER.request(
{
"cmd": "prepare",
"image": edited_frame,
"video": processed_video,
"area": RESOLUTION_AREA[resolution],
"prompt": prompt_text,
"out": prep_path,
},
ENCODER_PREPARE_TIMEOUT,
)
except EncoderTaskError as e:
print(f"[job] prepare failed: {e}", flush=True)
raise gr.Error("پردازش تصویر یا ویدیوی ورودی ناموفق بود؛ لطفاً فایل‌ها را بررسی کنید.")
except RuntimeError as e:
print(f"[job] encoder unavailable: {e}", flush=True)
raise gr.Error("انکودر موقتاً در دسترس نیست؛ چند لحظه بعد دوباره تلاش کنید.")
tensors = torch.load(prep_path, map_location="cpu")
height, width = prep["height"], prep["width"]
num_segments, target_frames = eng.segment_plan(real_frame_len)
with open(processed_video, "rb") as f:
video_bytes = f.read()
with open(edited_frame, "rb") as f:
img_bytes = f.read()
print(f"[job] res={resolution} frame={width}x{height} frames={real_frame_len} segments={num_segments}", flush=True)
return {
"format": STATE_FORMAT,
"model_repo": MODEL_REPO,
"resolution": resolution,
"height": height,
"width": width,
"crop_region": list(prep["crop_region"]),
"seed": random.randint(0, 2**31 - 1),
"prompt_text": prompt_text,
"prompt_embeds": tensors["prompt_embeds"],
"image_pixels": tensors["image_pixels"],
"clip_ref": tensors["clip_ref"],
"clip_drive": tensors["clip_drive"],
"video_bytes": video_bytes,
"img_bytes": img_bytes,
"audio_bytes": audio_bytes,
"real_frame_len": real_frame_len,
"num_segments": num_segments,
"seg_index": 0,
"cur_step": 0,
"latents": None,
"sched_state": None,
"drive_latents": None,
"prev_cond": None,
"tail_frames": None,
"segment_videos": [],
"frames_emitted": 0,
"timings": dict(DEFAULT_TIMINGS[resolution]),
}
def prepare_segment_pixels(state, video_path, segments, workdir):
res = ENCODER.request(
{
"cmd": "segments",
"video": video_path,
"real_frame_len": state["real_frame_len"],
"height": state["height"],
"width": state["width"],
"segments": [int(k) for k in segments],
"workdir": workdir,
},
ENCODER_SEGMENTS_TIMEOUT + ENCODER_SEGMENT_EXTRA_TIMEOUT * len(segments),
)
return {int(k): v for k, v in res["paths"].items()}
# --------------------------------------------------------------------------
# پنجره‌ی GPU: تا جایی که زمان اجازه دهد پیش می‌رود و وضعیت را برمی‌گرداند
# --------------------------------------------------------------------------
def _gpu_duration(resolution, job_path, window_seconds):
return int(window_seconds)
@spaces.GPU(duration=_gpu_duration, size=GPU_SIZE)
def run_gpu_window(resolution, job_path, window_seconds):
t_start = time.time()
try:
with torch.no_grad():
return _run_window(resolution, job_path, int(window_seconds), t_start)
finally:
vae.use_tiling = False
gc.collect()
torch.cuda.empty_cache()
_OOM_TYPES = tuple({getattr(torch, "OutOfMemoryError", torch.cuda.OutOfMemoryError), torch.cuda.OutOfMemoryError})
def _is_oom(err):
return isinstance(err, _OOM_TYPES) or "out of memory" in str(err).lower()
def _run_window(resolution, job_path, window_seconds, t_start):
device = torch.device("cuda")
bf16 = torch.bfloat16
gib = 1024 ** 3
job = torch.load(job_path, map_location="cpu", weights_only=False)
budget = window_seconds - GPU_SAFETY_SECONDS - WINDOW_INIT_ALLOWANCE
timings = dict(job["timings"])
result_path = job["result_path"]
progress_path = job.get("progress_path")
real_frame_len = int(job["real_frame_len"])
frames_emitted = int(job["frames_emitted"])
encoder_pool = ThreadPoolExecutor(max_workers=1)
pending_encodes = []
def report(text):
if not progress_path:
return
try:
with open(progress_path + ".tmp", "w", encoding="utf-8") as pf:
json.dump({"text": text}, pf, ensure_ascii=False)
os.replace(progress_path + ".tmp", progress_path)
except OSError:
pass
def left():
return budget - (time.time() - t_start)
def est(name):
return timings[name] * 1.15
def sync():
torch.cuda.synchronize()
height, width = job["height"], job["width"]
lat_h, lat_w = height // 8, width // 8
lat_frames = (eng.SEGMENT_FRAME_LENGTH - 1) // 4 + 1
grid_ref = torch.tensor([[lat_frames, lat_h // 2, lat_w // 2]], dtype=torch.long)
max_seq_len = (lat_frames + 1) * (lat_h // 2) * (lat_w // 2)
max_seq_len_ref = lat_frames * lat_h * lat_w // 4
token_bytes = max_seq_len * transformer.config.dim * 2
def reserve_bytes():
return 12 * token_bytes + int((8.0 + float(timings.get("mem_extra_gb", 0.0))) * gib)
prompt_embeds = job["prompt_embeds"].to(device, bf16)
negative_embeds = job["negative_embeds"].to(device, bf16)
prompt_ref_embeds = job["prompt_ref_embeds"].to(device, bf16)
clip_ref = job["clip_ref"].to(device, bf16)
clip_drive = job["clip_drive"].to(device, bf16)
ref_lat = encode_vae(vae, job["image_pixels"].to(device, torch.float32).unsqueeze(2))
mask_ref = get_i2v_mask(1, lat_h, lat_w, 1, device=device).to(ref_lat.dtype)
reference_image_latents = torch.cat([mask_ref, ref_lat[0]], dim=0)
del ref_lat, mask_ref
scheduler = copy.deepcopy(base_scheduler)
seg = job["seg_index"]
num_segments = job["num_segments"]
cur_step = job["cur_step"]
latents = job["latents"]
sched_state = job["sched_state"]
drive_latents = job["drive_latents"]
prev_cond = job["prev_cond"]
tail_frames = job["tail_frames"]
crop_top, crop_left, crop_h, crop_w = job["crop_region"]
workdir = job["workdir"]
completed = []
progressed = False
steps_this_window = 0
oom_count = 0
last_checkpoint = time.time()
def to_cpu(x, dtype=None):
if x is None:
return None
x = x.detach().to("cpu")
return x.to(dtype) if dtype is not None else x
def checkpoint():
nonlocal last_checkpoint
keep_drive = cur_step < NUM_INFERENCE_STEPS
result = {
"seg_index": seg,
"cur_step": cur_step,
"latents": to_cpu(latents, torch.float32),
"sched_state": sched_state,
"drive_latents": to_cpu(drive_latents) if keep_drive else None,
"prev_cond": to_cpu(prev_cond) if keep_drive else None,
"tail_frames": tail_frames,
"timings": dict(timings),
"completed": list(completed),
"frames_emitted": frames_emitted,
}
atomic_torch_save(result, result_path)
last_checkpoint = time.time()
def on_oom(stage, err):
nonlocal oom_count
oom_count += 1
timings["mem_extra_gb"] = float(timings.get("mem_extra_gb", 0.0)) + OOM_EXTRA_GB
print(f"[gpu] OOM در مرحله‌ی {stage} (بار {oom_count}) — افزایش حاشیه‌ی حافظه به "
f"{timings['mem_extra_gb']:.0f}GB", flush=True)
gc.collect()
torch.cuda.empty_cache()
if oom_count > MAX_OOM_RETRIES:
raise err
while seg < num_segments:
# ۱) انکود ویدیوی سگمنت + فریم‌های شرطی سگمنت قبل
if cur_step < NUM_INFERENCE_STEPS and (drive_latents is None or prev_cond is None):
if progressed and left() < est("encode") + est("extract") + CFG_PASSES * est("fwd"):
break
pix_path = job["segment_pixels"].get(seg)
if not pix_path:
break
t0 = time.time()
try:
report(f"بخش {seg + 1} از {num_segments} — آماده‌سازی ویدیوی حرکت")
pixels = eng.pixels_from_uint8(torch.load(pix_path, map_location="cpu"), device)
drive_latents = encode_vae(vae, pixels)
del pixels
prev_cond = eng.build_prev_cond_latents(
lambda x: encode_vae(vae, x), tail_frames, lat_h, lat_w, height, width, device
)
sync()
vae.use_tiling = False
except Exception as e:
if not _is_oom(e):
raise
pixels = None
drive_latents = None
prev_cond = None
on_oom("encode", e)
vae.enable_tiling()
continue
timings["encode"] = time.time() - t0
else:
if drive_latents is not None:
drive_latents = drive_latents.to(device)
if prev_cond is not None:
prev_cond = prev_cond.to(device)
# ۲) استخراج مرجع و گام‌های دینویز
if cur_step < NUM_INFERENCE_STEPS:
if progressed and left() < est("extract") + CFG_PASSES * est("fwd"):
break
reference_latents = torch.cat([reference_image_latents, prev_cond], dim=1).to(bf16)
scheduler.set_timesteps(NUM_INFERENCE_STEPS, device=device)
timesteps = scheduler.timesteps
if latents is None:
generator = torch.Generator(device=device).manual_seed(int(job["seed"]) + seg)
latents = randn_tensor(
(16, reference_latents.shape[1], lat_h, lat_w), generator=generator, device=device,
dtype=torch.float32,
)
sched_state = None
else:
latents = latents.to(device, torch.float32)
eng.load_scheduler_state(scheduler, sched_state, device)
t0 = time.time()
report(f"بخش {seg + 1} از {num_segments} — تحلیل حرکت مرجع")
kv_cache = WanAnimate2KVCache(transformer.config.num_layers)
try:
kv_policy.reset(reserve_bytes(), transformer.config.num_layers)
cond_mask = get_i2v_mask(
drive_latents.shape[2], lat_h, lat_w, eng.SEGMENT_FRAME_LENGTH, device=device
).to(drive_latents.dtype)
driving_condition = torch.cat([cond_mask, drive_latents[0]], dim=0)
del cond_mask
t_ref = torch.tensor([timesteps[0].item()], device=device, dtype=bf16)
transformer(
[drive_latents[0].to(bf16)],
timestep=t_ref,
encoder_hidden_states=[prompt_ref_embeds],
encoder_hidden_states_image=clip_drive,
condition_latents=[driving_condition.to(bf16)],
kv_cache=kv_cache,
kv_cache_mode="extract",
seq_len=max_seq_len_ref,
offset_grid_sizes=grid_ref,
)
del driving_condition
sync()
except Exception as e:
if not _is_oom(e):
raise
driving_condition = None
kv_cache.clear()
del kv_cache, reference_latents
sched_state = eng.scheduler_state_dict(scheduler)
on_oom("extract", e)
continue
timings["extract"] = time.time() - t0
print(f"[gpu] seg {seg + 1}/{num_segments} ref-cache layers: gpu={kv_policy.gpu_layers} "
f"cpu={kv_policy.cpu_layers} compact={kv_policy.x_layers}", flush=True)
oom_in_steps = None
while cur_step < NUM_INFERENCE_STEPS:
if steps_this_window > 0 and left() < CFG_PASSES * est("fwd"):
break
t0 = time.time()
t = timesteps[cur_step]
timestep = torch.stack([t])
latent_input = [latents.to(bf16)]
common = dict(
timestep=timestep,
condition_latents=[reference_latents],
kv_cache=kv_cache,
kv_cache_mode="cached",
seq_len=max_seq_len,
reference_grid_sizes=grid_ref,
origin_len=eng.SEGMENT_FRAME_LENGTH,
origin_area=[height, width],
encoder_hidden_states_image=clip_ref,
)
try:
noise_cond = transformer(latent_input, encoder_hidden_states=[prompt_embeds], is_uncondtion=False, **common).sample[0]
if CFG_PASSES == 2:
noise_uncond = transformer(latent_input, encoder_hidden_states=[negative_embeds], is_uncondtion=True, **common).sample[0]
noise_pred = noise_uncond + GUIDANCE_SCALE * (noise_cond - noise_uncond)
del noise_uncond
else:
noise_pred = noise_cond
del noise_cond
except Exception as e:
if not _is_oom(e):
raise
noise_cond = noise_uncond = noise_pred = None
latent_input = common = None
oom_in_steps = e
break
del latent_input
latents = scheduler.step(
noise_pred.unsqueeze(0), t, latents.unsqueeze(0), return_dict=False
)[0].squeeze(0)
del noise_pred
cur_step += 1
steps_this_window += 1
progressed = True
sync()
timings["fwd"] = (time.time() - t0) / CFG_PASSES
print(f"[gpu] seg {seg + 1}/{num_segments} step {cur_step}/{NUM_INFERENCE_STEPS} ({left():.0f}s left)", flush=True)
report(f"بخش {seg + 1} از {num_segments} — گام {cur_step} از {NUM_INFERENCE_STEPS}")
if cur_step < NUM_INFERENCE_STEPS and time.time() - last_checkpoint > CHECKPOINT_EVERY_SECONDS:
sched_state = eng.scheduler_state_dict(scheduler)
checkpoint()
kv_cache.clear()
del kv_cache, reference_latents
torch.cuda.empty_cache()
if cur_step < NUM_INFERENCE_STEPS:
sched_state = eng.scheduler_state_dict(scheduler)
if oom_in_steps is not None:
on_oom("denoise", oom_in_steps)
continue
break
# ۳) دیکود سگمنت کامل‌شده
if progressed and left() < est("decode"):
break
t0 = time.time()
try:
report(f"بخش {seg + 1} از {num_segments} — ساخت فریم‌ها")
latents = latents.to(device, torch.float32)
frames = decode_vae(vae, latents[:, 1:]) # [1, 3, 81, H, W]
except Exception as e:
if not _is_oom(e):
raise
frames = None
on_oom("decode", e)
vae.enable_tiling()
continue
if seg > 0:
frames = frames[:, :, eng.PREV_SEGMENT_COND_FRAMES:]
tail_frames = frames[0, :, -eng.PREV_SEGMENT_COND_FRAMES:].detach().to("cpu", bf16).clone()
video = frames[0, :, :, crop_top:crop_top + crop_h, crop_left:crop_left + crop_w]
video = ((video.float() / 2 + 0.5).clamp(0, 1) * 255).round().to(torch.uint8)
video = video.permute(1, 2, 3, 0).contiguous().cpu().numpy()
del frames
video = video[: max(0, real_frame_len - frames_emitted)]
frames_path = os.path.join(workdir, f"out_seg_{seg}.npy")
mp4_path = os.path.join(workdir, f"out_seg_{seg}.mp4")
if len(video) > 0:
np.save(frames_path, video)
pending_encodes.append(encoder_pool.submit(_npy_to_mp4, frames_path, mp4_path))
completed.append({"seg": seg, "path": frames_path, "mp4": mp4_path, "num_frames": int(video.shape[0])})
frames_emitted += int(video.shape[0])
del video
sync()
timings["decode"] = time.time() - t0
progressed = True
vae.use_tiling = False
seg += 1
cur_step = 0
latents = None
sched_state = None
drive_latents = None
prev_cond = None
torch.cuda.empty_cache()
checkpoint()
for fut in pending_encodes:
try:
fut.result(timeout=300)
except Exception as e:
print(f"[gpu] segment mp4 encode failed (parent will retry): {e}", flush=True)
encoder_pool.shutdown(wait=False)
checkpoint()
print(f"[gpu] window done in {time.time() - t_start:.1f}s — seg {seg}/{num_segments}, step {cur_step}", flush=True)
return result_path
# --------------------------------------------------------------------------
# حالت «جایگزینی شخصیت» (Mix / Replacement) — Wan2.2-Animate-14B
# --------------------------------------------------------------------------
MIX_STATE_FORMAT = "wan-animate-mix/v1"
MIX_RESOLUTIONS = ["480p", "720p"]
MIX_RESOLUTION_AREA = {"480p": 832 * 480, "720p": 1280 * 720}
MIX_DEFAULT_TIMINGS = {
"480p": {"prep_init": 45.0, "prep_frame": 0.07, "load": 110.0, "encode": 16.0, "fwd": 10.0, "decode": 20.0,
"align": mix.MIX_ALIGN_SECONDS},
"720p": {"prep_init": 45.0, "prep_frame": 0.10, "load": 110.0, "encode": 36.0, "fwd": 34.0, "decode": 45.0,
"align": mix.MIX_ALIGN_SECONDS},
}
MIX_WINDOW_OVERHEAD = {"480p": 10.0, "720p": 20.0} # انکود تصویر مرجع در شروع هر پنجره
MIX_MIN_WINDOW_SECONDS = {"480p": 300, "720p": 480}
MIX_ASSETS_WAIT_SECONDS = 3600
MIX_RESULT_KEYS = (
"prep_done", "kp2ds", "bboxes", "masks", "seg_index", "cur_step", "latents", "sched_state",
"cond_pose", "cond_y", "tail_frames", "timings", "ref_pixels", "clip_ref", "ref_src", "ref_aligned",
)
def _mix_need_align(state):
"""هم‌ترازی کادر مرجع فقط پیش از شروع تولید و فقط برای کارهایی که نسخه‌ی باکیفیت مرجع را دارند."""
return (
not state.get("ref_aligned", True)
and state.get("ref_src") is not None
and int(state.get("seg_index", 0)) == 0
and int(state.get("cur_step", 0)) == 0
and state.get("latents") is None
)
def is_mix_choice(value):
return "mix" in str(value or "").lower()
def normalize_mix_resolution(value):
return "720p" if "720" in str(value or "") else "480p"
def init_state_mix(input_video, edited_frame, resolution, prompt_text, workdir):
if not input_video:
raise gr.Error("لطفاً ویدیوی مرجع را آپلود کنید.")
if not edited_frame:
raise gr.Error("لطفاً تصویر مرجع را آپلود کنید.")
processed_video = os.path.join(workdir, "driving_30fps.mp4")
try:
transcode_to_model_fps(input_video, processed_video, fps=mix.MIX_FPS)
except Exception as e:
print(f"[mix-job] transcode failed: {e}", flush=True)
raise gr.Error("ویدیوی ورودی قابل خواندن نیست؛ لطفاً فایل ویدیوی دیگری امتحان کنید.")
real_frame_len = count_frames(processed_video)
if real_frame_len < 2:
raise gr.Error("ویدیوی ورودی فریم کافی ندارد.")
audio_path = os.path.join(workdir, "audio.wav")
audio_bytes = None
if extract_audio(input_video, audio_path):
with open(audio_path, "rb") as f:
audio_bytes = f.read()
prompt_text = (prompt_text or "").strip()
prep_path = os.path.join(workdir, "mix_prepared.pt")
try:
prep = ENCODER.request(
{
"cmd": "mix_prepare",
"image": edited_frame,
"video": processed_video,
"area": MIX_RESOLUTION_AREA[resolution],
"prompt": prompt_text,
"out": prep_path,
"clip_dir": MIX_ASSETS["model_dir"],
},
ENCODER_PREPARE_TIMEOUT,
)
except EncoderTaskError as e:
print(f"[mix-job] prepare failed: {e}", flush=True)
raise gr.Error("پردازش تصویر یا ویدیوی ورودی ناموفق بود؛ لطفاً فایل‌ها را بررسی کنید.")
except RuntimeError as e:
print(f"[mix-job] encoder unavailable: {e}", flush=True)
raise gr.Error("انکودر موقتاً در دسترس نیست؛ چند لحظه بعد دوباره تلاش کنید.")
tensors = torch.load(prep_path, map_location="cpu")
num_segments, _ = mix.mix_segment_plan(real_frame_len)
with open(processed_video, "rb") as f:
video_bytes = f.read()
print(f"[mix-job] res={resolution} frame={prep['width']}x{prep['height']} frames={real_frame_len} "
f"segments={num_segments}", flush=True)
return {
"format": MIX_STATE_FORMAT,
"model_repo": mix.MIX_MODEL_REPO,
"mode": "mix",
"resolution": resolution,
"height": int(prep["height"]),
"width": int(prep["width"]),
"seed": random.randint(0, 2**31 - 1),
"prompt_text": prompt_text,
"prompt_embeds": tensors["prompt_embeds"],
"clip_ref": tensors["clip_ref"],
"ref_pixels": tensors["ref_pixels"],
"ref_src": tensors.get("ref_src"),
"ref_aligned": tensors.get("ref_src") is None,
"video_bytes": video_bytes,
"audio_bytes": audio_bytes,
"real_frame_len": real_frame_len,
"num_segments": num_segments,
"prep_chunks": int(math.ceil(real_frame_len / mix.MIX_SAM_CHUNK)),
"prep_done": 0,
"kp2ds": [],
"bboxes": [],
"masks": [],
"seg_index": 0,
"cur_step": 0,
"latents": None,
"sched_state": None,
"cond_pose": None,
"cond_y": None,
"tail_frames": None,
"segment_videos": [],
"frames_emitted": 0,
"timings": dict(MIX_DEFAULT_TIMINGS[resolution]),
}
def _mix_min_window(state):
t = state["timings"]
res = state["resolution"]
need = 0.0
if state["prep_done"] < state["prep_chunks"]:
n = min(mix.MIX_SAM_CHUNK, state["real_frame_len"] - state["prep_done"] * mix.MIX_SAM_CHUNK)
need += t["prep_init"] + n * t["prep_frame"]
elif _mix_need_align(state):
need += t["prep_init"] + t.get("align", mix.MIX_ALIGN_SECONDS)
elif state["cur_step"] >= mix.MIX_STEPS:
need += t["decode"]
else:
need += t["load"] + t["fwd"] + (t["encode"] if state["cond_pose"] is None else 0.0)
need = need * 1.15 + MIX_WINDOW_OVERHEAD[res] + GPU_SAFETY_SECONDS + WINDOW_INIT_ALLOWANCE
return int(math.ceil(need))
def _mix_plan_window(state):
t = state["timings"]
res = state["resolution"]
steps = mix.MIX_STEPS
total = MIX_WINDOW_OVERHEAD[res]
remaining_prep = state["real_frame_len"] - state["prep_done"] * mix.MIX_SAM_CHUNK
if state["prep_done"] < state["prep_chunks"]:
total += t["prep_init"] + max(0, remaining_prep) * t["prep_frame"]
if _mix_need_align(state):
total += t.get("align", mix.MIX_ALIGN_SECONDS)
if state["prep_done"] >= state["prep_chunks"]:
total += t["prep_init"]
work_cap = (MAX_WINDOW_SECONDS - GPU_SAFETY_SECONDS - WINDOW_INIT_ALLOWANCE) / WINDOW_ESTIMATE_FACTOR
seg = state["seg_index"]
num_segments = state["num_segments"]
full_segment = t["encode"] + steps * t["fwd"] + t["decode"]
if seg < num_segments and total < work_cap:
if state["cur_step"] >= steps:
first = t["decode"]
else:
first = (steps - state["cur_step"]) * t["fwd"] + t["decode"]
first += t["encode"] if state["cond_pose"] is None else 0.0
if state["cur_step"] < steps or seg + 1 < num_segments:
total += t["load"]
total += first
for _ in range(seg + 1, num_segments):
if total + full_segment > work_cap:
break
total += full_segment
window = total * WINDOW_ESTIMATE_FACTOR + GPU_SAFETY_SECONDS + WINDOW_INIT_ALLOWANCE
window = int(math.ceil(max(MIX_MIN_WINDOW_SECONDS[res], min(MAX_WINDOW_SECONDS, window))))
return max(window, min(MAX_WINDOW_SECONDS, _mix_min_window(state)))
def _gpu_duration_mix(job_path, window_seconds):
return int(window_seconds)
@spaces.GPU(duration=_gpu_duration_mix, size=GPU_SIZE)
def run_mix_gpu_window(job_path, window_seconds):
t_start = time.time()
try:
with torch.no_grad():
return _run_mix_window(job_path, int(window_seconds), t_start)
finally:
vae.use_tiling = False
gc.collect()
torch.cuda.empty_cache()
def _run_mix_window(job_path, window_seconds, t_start):
device = torch.device("cuda")
bf16 = torch.bfloat16
job = torch.load(job_path, map_location="cpu", weights_only=False)
budget = window_seconds - GPU_SAFETY_SECONDS - WINDOW_INIT_ALLOWANCE
timings = dict(job["timings"])
result_path = job["result_path"]
progress_path = job.get("progress_path")
workdir = job["workdir"]
height, width = job["height"], job["width"]
lat_h, lat_w = height // 8, width // 8
real_frame_len = int(job["real_frame_len"])
frames_emitted = int(job["frames_emitted"])
chunk = mix.MIX_SAM_CHUNK
steps = mix.MIX_STEPS
prep_chunks = int(job["prep_chunks"])
prep_done = int(job["prep_done"])
kp_list = list(job["kp2ds"])
box_list = list(job["bboxes"])
mask_list = list(job["masks"])
seg = int(job["seg_index"])
num_segments = int(job["num_segments"])
cur_step = int(job["cur_step"])
latents = job["latents"]
sched_state = job["sched_state"]
cond_pose = job["cond_pose"]
cond_y = job["cond_y"]
tail_frames = job["tail_frames"]
ref_pixels = job["ref_pixels"]
clip_ref_cpu = job["clip_ref"]
ref_src = job.get("ref_src")
ref_aligned = bool(job.get("ref_aligned", True))
mix_vae = None
def need_align():
return (not ref_aligned) and ref_src is not None and seg == 0 and cur_step == 0 and latents is None
completed = []
progressed = False
steps_this_window = 0
last_checkpoint = time.time()
encoder_pool = ThreadPoolExecutor(max_workers=1)
pending_encodes = []
frames_src = mix.FrameSource(job["video_path"], height, width, real_frame_len)
def report(text):
if not progress_path:
return
try:
with open(progress_path + ".tmp", "w", encoding="utf-8") as pf:
json.dump({"text": text}, pf, ensure_ascii=False)
os.replace(progress_path + ".tmp", progress_path)
except OSError:
pass
def left():
return budget - (time.time() - t_start)
def est(name):
return timings[name] * 1.15
def to_cpu(x, dtype=None):
if x is None:
return None
x = x.detach().to("cpu")
return x.to(dtype) if dtype is not None else x
def checkpoint():
nonlocal last_checkpoint
keep_cond = cur_step < steps
atomic_torch_save({
"prep_done": prep_done,
"kp2ds": list(kp_list),
"bboxes": list(box_list),
"masks": list(mask_list),
"seg_index": seg,
"cur_step": cur_step,
"latents": to_cpu(latents, torch.float32),
"sched_state": sched_state,
"cond_pose": to_cpu(cond_pose, bf16) if keep_cond else None,
"cond_y": to_cpu(cond_y, bf16) if keep_cond else None,
"tail_frames": tail_frames,
"timings": dict(timings),
"ref_pixels": ref_pixels,
"clip_ref": clip_ref_cpu,
"ref_src": ref_src,
"ref_aligned": ref_aligned,
"completed": list(completed),
"frames_emitted": frames_emitted,
}, result_path)
last_checkpoint = time.time()
def get_vae():
nonlocal mix_vae
if mix_vae is None:
mix_vae, is_fp32 = mix.load_mix_vae(job["model_dir"], device, vae)
print(f"[mix-gpu] VAE: {'float32 (Animate-14B)' if is_fp32 else 'shared bf16'}", flush=True)
return mix_vae
def encode_fn(x):
v = get_vae()
try:
return encode_vae(v, x)
except Exception as e:
if not _is_oom(e):
raise
gc.collect()
torch.cuda.empty_cache()
v.enable_tiling()
return encode_vae(v, x)
# ---------------- ۱) پیش‌پردازش: نقاط بدن + ماسک شخصیت ----------------
if prep_done < prep_chunks or need_align():
t0 = time.time()
report("آماده‌سازی تحلیل بدن و ماسک شخصیت")
pose_models = mix.load_pose_models(job["aux_dir"])
sam = mix.build_sam_predictor(job["aux_dir"], device) if prep_done < prep_chunks else None
timings["prep_init"] = time.time() - t0
while prep_done < prep_chunks:
start = prep_done * chunk
end = min(real_frame_len, start + chunk)
n = end - start
if progressed and left() < n * est("prep_frame"):
break
t0 = time.time()
report(f"تحلیل حرکت و ماسک شخصیت — بخش {prep_done + 1} از {prep_chunks}")
frames = frames_src.get(list(range(start, end)))
kp = mix.estimate_kp2ds(pose_models, frames)
all_kp = np.concatenate([k.numpy() for k in kp_list] + [kp], 0)
metas = mix.metas_from_kp2ds(all_kp, width, height)[start:end]
boxes = mix.face_bboxes(metas, height, width)
masks = mix.sam_chunk_masks(sam, frames, metas, device)
aug = [mix.augment_mask(m) for m in masks]
kp_list.append(torch.from_numpy(kp))
box_list.append(torch.from_numpy(boxes))
mask_list.append(mix.pack_masks(aug))
prep_done += 1
progressed = True
timings["prep_frame"] = (time.time() - t0) / max(1, n)
print(f"[mix-gpu] prep chunk {prep_done}/{prep_chunks} ({n} frames, "
f"{timings['prep_frame']:.3f}s/frame, {left():.0f}s left)", flush=True)
del frames, masks, aug, all_kp, metas
checkpoint()
# هم‌ترازی کادر تصویر مرجع با صورت شخص داخل ویدیو (یک بار، پیش از شروع تولید)
if prep_done >= prep_chunks and need_align() and not (
progressed and left() < timings.get("align", mix.MIX_ALIGN_SECONDS) * 1.15
):
t0 = time.time()
report("هم‌ترازی چهره‌ی تصویر مرجع با ویدیو")
try:
drv_kp = np.concatenate([k.numpy() for k in kp_list], 0)
src = ref_src.numpy() if isinstance(ref_src, torch.Tensor) else np.asarray(ref_src)
aligned, info = mix.align_reference(pose_models, src.astype(np.uint8), drv_kp, height, width)
if aligned is not None:
new_clip = mix.clip_encode_on_gpu(job["model_dir"], aligned, device)
ref_pixels = torch.from_numpy(np.ascontiguousarray(aligned)).permute(2, 0, 1).contiguous()
clip_ref_cpu = new_clip
print(f"[mix-gpu] reference aligned: {info}", flush=True)
else:
print(f"[mix-gpu] reference alignment skipped ({info}); using padded reference", flush=True)
except Exception as e:
print(f"[mix-gpu] reference alignment failed ({type(e).__name__}: {e}); "
f"using padded reference", flush=True)
gc.collect()
torch.cuda.empty_cache()
ref_aligned = True
ref_src = None
progressed = True
timings["align"] = time.time() - t0
checkpoint()
del pose_models, sam
gc.collect()
torch.cuda.empty_cache()
# ---------------- ۲) تولید سگمنت‌ها ----------------
mix_tf = None
ref_cond = None
scheduler = copy.deepcopy(MIX_ASSETS["scheduler"])
prompt_embeds = job["prompt_embeds"].to(device, bf16).unsqueeze(0)
clip_ref = None
kp_all = np.concatenate([k.numpy() for k in kp_list], 0) if prep_done >= prep_chunks else None
box_all = np.concatenate([b.numpy() for b in box_list], 0) if prep_done >= prep_chunks else None
while prep_done >= prep_chunks and not need_align() and seg < num_segments:
if cur_step < steps:
if mix_tf is None:
if progressed and left() < est("load") + est("fwd") + (est("encode") if cond_pose is None else 0):
break
t0 = time.time()
report("بارگذاری مدل روی کارت گرافیک")
mix_tf = mix.load_mix_transformer(job["model_dir"], job["aux_dir"], device)
timings["load"] = time.time() - t0
print(f"[mix-gpu] transformer loaded in {timings['load']:.1f}s", flush=True)
if progressed and left() < est("fwd") + (est("encode") if cond_pose is None else 0):
break
if ref_cond is None:
ref_cond = mix.reference_condition(encode_fn, ref_pixels, device).to(bf16)
if clip_ref is None:
clip_ref = clip_ref_cpu.to(device, bf16)
t0 = time.time()
report(f"بخش {seg + 1} از {num_segments} — آماده‌سازی شرط‌های ویدیو")
pixels = mix.segment_pixels(frames_src, seg, real_frame_len, height, width, kp_all, box_all, mask_list)
face = mix.u8_to_pm1(pixels["face"], device, bf16)
if cond_pose is None:
cond_pose, cond_y = mix.segment_conditions(
encode_fn, pixels, tail_frames, seg == 0, height, width, device
)
torch.cuda.synchronize()
timings["encode"] = time.time() - t0
get_vae().use_tiling = False
else:
cond_pose = cond_pose.to(device)
cond_y = cond_y.to(device)
del pixels
y = torch.cat([ref_cond, cond_y.to(bf16)], dim=1).unsqueeze(0)
pose_lat = cond_pose.to(bf16)
scheduler.set_timesteps(steps, device=device)
timesteps = scheduler.timesteps
if latents is None:
generator = torch.Generator(device=device).manual_seed(int(job["seed"]) + seg)
latents = randn_tensor(
(1, 16, mix.MIX_LATENT_FRAMES + 1, lat_h, lat_w), generator=generator,
device=device, dtype=torch.float32,
)
sched_state = None
else:
latents = latents.to(device, torch.float32)
eng.load_scheduler_state(scheduler, sched_state, device)
torch.cuda.empty_cache()
while cur_step < steps:
if steps_this_window > 0 and left() < est("fwd"):
break
t0 = time.time()
t = timesteps[cur_step]
noise_pred = mix_tf(
hidden_states=torch.cat([latents.to(bf16), y], dim=1),
timestep=t.expand(1),
encoder_hidden_states=prompt_embeds,
encoder_hidden_states_image=clip_ref,
pose_hidden_states=pose_lat,
face_pixel_values=face,
return_dict=False,
)[0]
latents = scheduler.step(noise_pred.float(), t, latents, return_dict=False)[0]
del noise_pred
cur_step += 1
steps_this_window += 1
progressed = True
torch.cuda.synchronize()
timings["fwd"] = time.time() - t0
print(f"[mix-gpu] seg {seg + 1}/{num_segments} step {cur_step}/{steps} ({left():.0f}s left)", flush=True)
report(f"بخش {seg + 1} از {num_segments} — گام {cur_step} از {steps}")
if cur_step < steps and time.time() - last_checkpoint > CHECKPOINT_EVERY_SECONDS:
sched_state = eng.scheduler_state_dict(scheduler)
checkpoint()
del face, y, pose_lat
if cur_step < steps:
sched_state = eng.scheduler_state_dict(scheduler)
break
# دیکود سگمنت کامل‌شده
if progressed and left() < est("decode"):
break
t0 = time.time()
report(f"بخش {seg + 1} از {num_segments} — ساخت فریم‌ها")
latents = latents.to(device, torch.float32)
dec_vae = get_vae()
try:
frames = decode_vae(dec_vae, latents[:, :, 1:])
except Exception as e:
if not _is_oom(e):
raise
gc.collect()
torch.cuda.empty_cache()
dec_vae.enable_tiling()
frames = decode_vae(dec_vae, latents[:, :, 1:])
if seg > 0:
frames = frames[:, :, mix.MIX_PREV_FRAMES:]
tail_frames = frames[0, :, -mix.MIX_PREV_FRAMES:].detach().to("cpu", bf16).clone()
video = ((frames[0].float() / 2 + 0.5).clamp(0, 1) * 255).round().to(torch.uint8)
video = video.permute(1, 2, 3, 0).contiguous().cpu().numpy()
del frames
video = video[: max(0, real_frame_len - frames_emitted)]
frames_path = os.path.join(workdir, f"mix_seg_{seg}.npy")
mp4_path = os.path.join(workdir, f"mix_seg_{seg}.mp4")
if len(video) > 0:
np.save(frames_path, video)
pending_encodes.append(encoder_pool.submit(_npy_to_mp4, frames_path, mp4_path, mix.MIX_FPS))
completed.append({"seg": seg, "path": frames_path, "mp4": mp4_path, "num_frames": int(video.shape[0])})
frames_emitted += int(video.shape[0])
del video
torch.cuda.synchronize()
timings["decode"] = time.time() - t0
dec_vae.use_tiling = False
progressed = True
seg += 1
cur_step = 0
latents = None
sched_state = None
cond_pose = None
cond_y = None
torch.cuda.empty_cache()
checkpoint()
for fut in pending_encodes:
try:
fut.result(timeout=300)
except Exception as e:
print(f"[mix-gpu] segment mp4 encode failed (parent will retry): {e}", flush=True)
encoder_pool.shutdown(wait=False)
frames_src.close()
del mix_tf
if mix_vae is not None and mix_vae is not vae:
mix_vae = None
checkpoint()
print(f"[mix-gpu] window done in {time.time() - t_start:.1f}s — prep {prep_done}/{prep_chunks}, "
f"seg {seg}/{num_segments}, step {cur_step}", flush=True)
return result_path
def _wait_mix_assets():
"""(generator) تا آماده شدن فایل‌های حالت جایگزینی صبر می‌کند؛ True/False برمی‌گرداند."""
started = time.time()
while not _mix_assets_ready.wait(PROGRESS_POLL_SECONDS):
if time.time() - started > MIX_ASSETS_WAIT_SECONDS:
return False
yield None, None, (f"⏳ **آماده‌سازی مدل جایگزینی شخصیت روی سرور...** "
f"({_format_seconds(time.time() - started)})")
if MIX_ASSETS["scheduler"] is None:
yield None, None, "⏳ **تلاش دوباره برای آماده‌سازی مدل جایگزینی شخصیت...**"
return _download_mix_assets()
return True
def _animate_mix(state, input_video, edited_frame, resolution_choice, prompt_text, workdir):
ok = yield from _wait_mix_assets()
if not ok:
raise gr.Error("مدل جایگزینی شخصیت موقتاً در دسترس نیست؛ چند دقیقه بعد دوباره تلاش کنید.")
if state is None:
yield None, None, "⏳ **در حال آماده‌سازی ویدیو، تصویر و انکودرها...**"
state = init_state_mix(input_video, edited_frame, normalize_mix_resolution(resolution_choice),
prompt_text, workdir)
else:
yield None, None, "⏳ **فایل وضعیت بازیابی شد؛ ادامه‌ی رندر از آخرین نقطه...**"
fixed = ENCODER.fixed_embeds()
# فایل‌های وضعیت قدیمی (پیش از هم‌ترازی مرجع) بدون تغییر ادامه می‌یابند
state.setdefault("ref_src", None)
state.setdefault("ref_aligned", state["ref_src"] is None)
resolution = state["resolution"]
num_segments = state["num_segments"]
video_path = write_bytes(state["video_bytes"], ".mp4", workdir)
window = _mix_plan_window(state)
result_path = os.path.join(workdir, "result.pt")
progress_path = os.path.join(workdir, "progress.json")
job = {
"workdir": workdir,
"result_path": result_path,
"progress_path": progress_path,
"video_path": video_path,
"model_dir": MIX_ASSETS["model_dir"],
"aux_dir": MIX_ASSETS["aux_dir"],
"prompt_embeds": state["prompt_embeds"] if state["prompt_embeds"] is not None else fixed["mix"],
}
for key in ("height", "width", "seed", "clip_ref", "ref_pixels", "real_frame_len", "num_segments",
"prep_chunks", "frames_emitted") + MIX_RESULT_KEYS:
job[key] = state[key]
job_path = os.path.join(workdir, "job.pt")
torch.save(job, job_path)
del job
header = (f"🔄 **در حال رندر جایگزینی شخصیت ({resolution}) روی کارت کامل — "
f"{num_segments} بخش، مهلت این مرحله تا {_format_seconds(window)} دقیقه**")
yield None, None, header
result, failure = yield from _run_gpu_with_recovery(
_mix_min_window(state), run_mix_gpu_window, lambda w: (job_path, w),
result_path, window, progress_path, header,
)
if result is not None:
_apply_result(state, result, workdir, keys=MIX_RESULT_KEYS, fps=mix.MIX_FPS)
if state["seg_index"] < num_segments:
state_path = _save_state_file(state)
if state["prep_done"] < state["prep_chunks"]:
progress_line = (f"پیشرفت: تحلیل حرکت و ماسک — بخش **{state['prep_done']}** از "
f"**{state['prep_chunks']}** کامل شده.\n\n")
else:
progress_line = (
f"پیشرفت: بخش **{state['seg_index']}** از **{num_segments}** کامل شده؛ "
f"بخش جاری در گام **{state['cur_step']}** از **{mix.MIX_STEPS}** است.\n\n"
)
if failure is None:
status_msg = "⚠️ **این مرحله از رندر ذخیره شد.**\n\n" + progress_line + _CONTINUE_STEPS
else:
status_msg = (
"⚠️ **پنجره‌ی GPU این بار کامل انجام نشد، اما پیشرفت کار ذخیره شد.**\n\n"
+ progress_line
+ f"علت (از سمت ZeroGPU): `{failure[:400]}`\n\n"
+ _CONTINUE_STEPS
)
_round_sink.key = _progress_key(state)
yield None, state_path, status_msg
return
final_path = _finalize_output(state, workdir, "wan_mix")
_round_sink.key = "final"
yield final_path, None, "✅ **رندر کل ویدیو با موفقیت به پایان رسید!**"
# --------------------------------------------------------------------------
# تابع اصلی (API: animate_scene)
# --------------------------------------------------------------------------
def do_transfer(out_file):
if not out_file:
raise gr.Error("فایل وضعیتی برای انتقال یافت نشد! ابتدا اجازه دهید فرآیند شروع شده و فایل وضعیت صادر شود.")
msg = ("🔄 **فایل وضعیت با موفقیت به بخش ورودی منتقل شد.**\n\n"
"حالا می‌توانید آی‌پی خود را تغییر داده و مجدداً روی دکمه **Wan Animate 🦆** کلیک کنید تا رندر ادامه یابد.")
return out_file, msg
def _load_state(state_file):
if not state_file or not os.path.exists(state_file):
return None
try:
state = torch.load(state_file, map_location="cpu", weights_only=True)
except Exception as e:
print(f"[resume] خواندن فایل وضعیت ناموفق بود: {e}", flush=True)
return None
if isinstance(state, dict) and state.get("format") == MIX_STATE_FORMAT and state.get("model_repo") == mix.MIX_MODEL_REPO:
return state
if not isinstance(state, dict) or state.get("format") != STATE_FORMAT or state.get("model_repo") != MODEL_REPO:
print("[resume] فایل وضعیت مربوط به مدل/نسخه‌ی قبلی است؛ کار از ابتدا با مدل جدید شروع می‌شود.", flush=True)
return None
return state
def _min_window_seconds(state):
t = state["timings"]
res = state["resolution"]
if state["cur_step"] >= NUM_INFERENCE_STEPS:
need = t["decode"] * 1.15
else:
need = t["extract"] * 1.15 + CFG_PASSES * t["fwd"] * 1.15
if state["drive_latents"] is None:
need += t["encode"] * 1.15
need += WINDOW_OVERHEAD_SECONDS[res] + GPU_SAFETY_SECONDS + WINDOW_INIT_ALLOWANCE
return int(math.ceil(need))
def _plan_window(state, video_frame_bytes):
"""
مدت پنجره را از روی کار باقی‌مانده تخمین می‌زند (حداکثر ۲۰ دقیقه) و فهرست بخش‌هایی را
برمی‌گرداند که پیکسل‌هایشان باید قبل از پنجره آماده شوند.
"""
t = state["timings"]
res = state["resolution"]
n_steps = NUM_INFERENCE_STEPS
seg = state["seg_index"]
num_segments = state["num_segments"]
full_segment = t["encode"] + t["extract"] + n_steps * CFG_PASSES * t["fwd"] + t["decode"]
needed = []
if state["cur_step"] >= n_steps:
first = t["decode"]
else:
first = t["extract"] + (n_steps - state["cur_step"]) * CFG_PASSES * t["fwd"] + t["decode"]
if state["drive_latents"] is None:
first += t["encode"]
needed.append(seg)
total = WINDOW_OVERHEAD_SECONDS[res] + first
work_cap = (MAX_WINDOW_SECONDS - GPU_SAFETY_SECONDS - WINDOW_INIT_ALLOWANCE) / WINDOW_ESTIMATE_FACTOR
pixel_bytes = len(needed) * video_frame_bytes
try:
free_disk = shutil.disk_usage(tempfile.gettempdir()).free - DISK_MARGIN_BYTES
except OSError:
free_disk = MAX_PIXEL_DISK_BYTES
disk_cap = max(video_frame_bytes, min(MAX_PIXEL_DISK_BYTES, free_disk))
for k in range(seg + 1, num_segments):
if total + full_segment > work_cap or pixel_bytes + video_frame_bytes > disk_cap:
break
total += full_segment
needed.append(k)
pixel_bytes += video_frame_bytes
window = total * WINDOW_ESTIMATE_FACTOR + GPU_SAFETY_SECONDS + WINDOW_INIT_ALLOWANCE
window = int(math.ceil(max(MIN_WINDOW_SECONDS[res], min(MAX_WINDOW_SECONDS, window))))
return max(window, min(MAX_WINDOW_SECONDS, _min_window_seconds(state))), needed
def _error_text(err):
title = str(getattr(err, "title", "") or "")
message = str(getattr(err, "message", "") or err)
return title, message
def _apply_result(state, result, workdir, keys=None, fps=None):
for item in sorted(result["completed"], key=lambda x: x["seg"]):
if item["seg"] < state["seg_index"] or item["num_frames"] <= 0:
continue
mp4_path = item["mp4"]
if not os.path.exists(mp4_path) and os.path.exists(item["path"]):
_npy_to_mp4(item["path"], mp4_path, fps)
if not os.path.exists(mp4_path):
print(f"[result] فایل بخش {item['seg']} پیدا نشد", flush=True)
continue
with open(mp4_path, "rb") as f:
state["segment_videos"].append(f.read())
os.remove(mp4_path)
state["frames_emitted"] = int(result.get("frames_emitted", state["frames_emitted"]))
keys = keys or ("seg_index", "cur_step", "latents", "sched_state", "drive_latents", "prev_cond", "tail_frames", "timings")
for key in keys:
state[key] = result[key]
def _save_state_file(state):
state_path = os.path.join(OUTPUT_DIR, f"wan_state_{uuid.uuid4().hex[:6]}.pt")
atomic_torch_save(state, state_path)
return state_path
_CONTINUE_STEPS = (
"👇 **ادامه‌ی کار:**\n"
"۱. روی دکمه **انتقال سریع فایل وضعیت به بخش ورودی 🔄** کلیک کنید.\n"
"۲. در صورت نیاز آی‌پی خود را تغییر دهید تا سهمیه تمدید شود.\n"
"۳. مجدداً روی دکمه **Wan Animate 🦆** کلیک کنید."
)
def _format_seconds(sec):
sec = int(max(0, sec))
return f"{sec // 60}:{sec % 60:02d}"
def _call_gpu_streaming(gpu_fn, gpu_args, progress_path, header):
"""
پنجره‌ی GPU را در یک رشته‌ی جدا (با همان context درخواست Gradio، تا سهمیه‌ی همان کاربر
استفاده شود) اجرا می‌کند و هر چند ثانیه وضعیت پیشرفت را yield می‌کند.
"""
ctx = contextvars.copy_context()
executor = ThreadPoolExecutor(max_workers=1)
future = executor.submit(ctx.run, gpu_fn, *gpu_args)
started = time.time()
try:
while True:
done, _ = futures_wait([future], timeout=PROGRESS_POLL_SECONDS)
if done:
return future.result()
detail = "در صف GPU یا در حال بارگذاری مدل روی کارت..."
try:
with open(progress_path, "r", encoding="utf-8") as pf:
detail = json.load(pf).get("text") or detail
except (OSError, ValueError):
pass
yield None, None, f"{header}\n\n⏱️ {_format_seconds(time.time() - started)} — {detail}"
finally:
executor.shutdown(wait=False)
def _run_gpu_with_recovery(min_window, gpu_fn, make_args, result_path, window, progress_path, header):
"""
(generator) پنجره‌ی GPU را اجرا می‌کند، پیشرفت را yield می‌کند و هیچ‌وقت خطا پرتاب نمی‌کند.
مقدار بازگشتی: (result یا None، متن علت شکست یا None)
"""
min_window = min(MAX_WINDOW_SECONDS, min_window)
other_retries = 1
shrink_retries = 6
last_reason = None
while True:
try:
yield from _call_gpu_streaming(gpu_fn, make_args(window), progress_path, header)
return torch.load(result_path, map_location="cpu", weights_only=False), None
except Exception as err:
title, message = _error_text(err)
print(f"[gpu-call] window={window}s failed: {type(err).__name__}: {title} {message}", flush=True)
last_reason = f"{title}: {message}" if title else message
# اگر پنجره قبل از قطع شدن بخشی از کار را ذخیره کرده، همان را نگه می‌داریم
if os.path.exists(result_path):
try:
return torch.load(result_path, map_location="cpu", weights_only=False), last_reason
except Exception as e:
print(f"[gpu-call] checkpoint unreadable: {e}", flush=True)
low = f"{title} {message}".lower()
if "illegal duration" in low or "larger than the maximum allowed" in low:
new_window = int(window * 2 // 3)
if shrink_retries > 0 and new_window >= min_window:
shrink_retries -= 1
window = new_window
continue
return None, last_reason
if "quota" in low or "credits" in low or "runs limit" in low or "gpu limit" in low:
m = re.search(r"\((\d+)s requested vs\. (\d+)s left\)", message)
if m and shrink_retries > 0:
left_s = int(m.group(2))
factor = 2 if GPU_SIZE == "xlarge" else 1
new_window = left_s // factor - 2
if min_window <= new_window < window:
shrink_retries -= 1
window = new_window
continue
return None, last_reason
if other_retries > 0:
other_retries -= 1
time.sleep(3)
continue
return None, last_reason
# --------------------------------------------------------------------------
# هر «مرحله‌ی رندر» در پس‌زمینه اجرا می‌شود و به اتصال درخواست وابسته نیست.
# - اگر اتصال کلاینت وسط کار قطع شود، پنجره‌ی GPU ادامه می‌یابد و نتیجه نگه داشته می‌شود.
# - اگر همان درخواست (همان ورودی یا همان فایل وضعیت) دوباره برسد، به همان رندر در حال اجرا
# وصل می‌شود یا نتیجه‌ی آماده‌اش را می‌گیرد؛ پس دو پنجره‌ی GPU برای یک کار همزمان اجرا
# نمی‌شود و سهمیه دوبار مصرف نمی‌شود.
# --------------------------------------------------------------------------
RUN_CACHE_SECONDS = 45 * 60
_round_sink = threading.local()
class RenderRound:
def __init__(self, key):
self.key = key
self.done = threading.Event()
self.progress = "⏳ **در حال شروع پردازش...**"
self.outputs = None
self.error = None
self.result_key = None
self.finished_at = None
_ROUNDS = {}
_ROUNDS_LOCK = threading.Lock()
def _file_digest(path):
h = hashlib.sha1()
with open(path, "rb") as f:
for chunk in iter(lambda: f.read(4 * 1024 * 1024), b""):
h.update(chunk)
return h.hexdigest()
def _state_job_id(state):
return f"{state.get('seed')}-{state.get('real_frame_len')}-{state.get('height')}x{state.get('width')}"
def _progress_key(state):
mode = "mix" if state.get("format") == MIX_STATE_FORMAT else "motion"
return (f"state:{mode}:{_state_job_id(state)}:{state['seg_index']}:{state['cur_step']}:"
f"{state.get('prep_done', 0)}:{len(state.get('segment_videos') or [])}")
def _round_reusable(rnd):
"""رندر تمام‌شده فقط وقتی دوباره تحویل داده می‌شود که واقعاً پیشرفتی داشته باشد."""
if rnd.error is not None or rnd.outputs is None:
return False
if time.time() - (rnd.finished_at or 0) > RUN_CACHE_SECONDS:
return False
return rnd.result_key is not None and rnd.result_key != rnd.key
def _prune_rounds():
now = time.time()
with _ROUNDS_LOCK:
for key in [k for k, r in _ROUNDS.items()
if r.done.is_set() and now - (r.finished_at or now) > RUN_CACHE_SECONDS]:
del _ROUNDS[key]
def _execute_round(rnd, mode, state, input_video, edited_frame, resolution_choice, prompt_text):
workdir = tempfile.mkdtemp(prefix=WORKDIR_PREFIX)
_round_sink.key = None
try:
if mode == "mix":
gen = _animate_mix(state, input_video, edited_frame, resolution_choice, prompt_text, workdir)
else:
gen = _animate_motion(state, input_video, edited_frame, resolution_choice, prompt_text, workdir)
last = None
for video, state_path, message in gen:
if video is None and state_path is None:
rnd.progress = message
else:
last = (video, state_path, message)
rnd.outputs = last
rnd.result_key = _round_sink.key
if last is None:
rnd.error = "پاسخی از رندر دریافت نشد؛ دوباره تلاش کنید."
except gr.Error as e:
rnd.error = str(getattr(e, "message", "") or e)
except Exception as e:
print(f"[round] unexpected error: {type(e).__name__}: {e}", flush=True)
traceback.print_exc()
rnd.error = "خطای غیرمنتظره در رندر رخ داد؛ چند لحظه بعد دوباره تلاش کنید."
finally:
shutil.rmtree(workdir, ignore_errors=True)
rnd.finished_at = time.time()
rnd.done.set()
def animate_scene(input_video, edited_frame, resolution_choice, prompt_text="", state_file=None,
progress=gr.Progress(track_tqdm=False)):
"""
resolution_choice:
"480p" / "720p" → انتقال حرکت (Wan2.2-Animate-2 Distilled)
"mix-480p" / "mix-720p" → جایگزینی شخصیت در ویدیو (Wan2.2-Animate-14B, replace)
"""
prune_old_files()
_prune_rounds()
state = _load_state(state_file)
if state is not None:
mode = "mix" if state.get("format") == MIX_STATE_FORMAT else "motion"
key = _progress_key(state)
else:
if not input_video:
raise gr.Error("لطفاً ویدیوی مرجع را آپلود کنید.")
if not edited_frame:
raise gr.Error("لطفاً تصویر مرجع را آپلود کنید.")
mode = "mix" if is_mix_choice(resolution_choice) else "motion"
res = normalize_mix_resolution(resolution_choice) if mode == "mix" else normalize_resolution(resolution_choice)
prompt_hash = hashlib.sha1((prompt_text or "").strip().encode("utf-8")).hexdigest()[:12]
key = f"new:{mode}:{res}:{_file_digest(input_video)}:{_file_digest(edited_frame)}:{prompt_hash}"
with _ROUNDS_LOCK:
rnd = _ROUNDS.get(key)
start = rnd is None or (rnd.done.is_set() and not _round_reusable(rnd))
if start:
rnd = RenderRound(key)
_ROUNDS[key] = rnd
ctx = contextvars.copy_context()
threading.Thread(
target=ctx.run,
args=(_execute_round, rnd, mode, state, input_video, edited_frame, resolution_choice, prompt_text),
daemon=True,
).start()
if not start:
print(f"[round] attached to existing render ({'finished' if rnd.done.is_set() else 'running'}): {key[:80]}",
flush=True)
yield None, None, "🔗 **این مرحله از قبل در حال رندر است؛ به همان رندر وصل شدیم...**"
started = time.time()
while not rnd.done.wait(PROGRESS_POLL_SECONDS):
yield None, None, rnd.progress
if time.time() - started > 6 * 3600:
break
if rnd.error is not None:
raise gr.Error(rnd.error)
if rnd.outputs is None:
raise gr.Error("پاسخی از رندر دریافت نشد؛ دوباره تلاش کنید.")
yield rnd.outputs
def _animate_motion(state, input_video, edited_frame, resolution_choice, prompt_text, workdir):
if state is None:
yield None, None, "⏳ **در حال آماده‌سازی ویدیو، تصویر و انکودرها...**"
state = init_state(input_video, edited_frame, normalize_resolution(resolution_choice), prompt_text, workdir)
else:
yield None, None, "⏳ **فایل وضعیت بازیابی شد؛ ادامه‌ی رندر از آخرین نقطه...**"
fixed = ENCODER.fixed_embeds()
resolution = state["resolution"]
video_path = write_bytes(state["video_bytes"], ".mp4", workdir)
seg = state["seg_index"]
num_segments = state["num_segments"]
# مدت پنجره و بخش‌هایی که در این پنجره به آن‌ها می‌رسیم
frame_bytes = 3 * eng.SEGMENT_FRAME_LENGTH * state["height"] * state["width"]
window, needed = _plan_window(state, frame_bytes)
if needed:
yield None, None, f"⏳ **آماده‌سازی فریم‌های {len(needed)} بخش از ویدیو...**"
try:
segment_pixels = prepare_segment_pixels(state, video_path, needed, workdir) if needed else {}
except Exception as e:
print(f"[job] segment preparation failed: {e}", flush=True)
state_path = _save_state_file(state)
_round_sink.key = _progress_key(state)
yield None, state_path, (
"⚠️ **آماده‌سازی فریم‌ها این بار انجام نشد، اما کار ذخیره شد.**\n\n" + _CONTINUE_STEPS
)
return
result_path = os.path.join(workdir, "result.pt")
progress_path = os.path.join(workdir, "progress.json")
job = {
"workdir": workdir,
"result_path": result_path,
"progress_path": progress_path,
"real_frame_len": state["real_frame_len"],
"frames_emitted": state["frames_emitted"],
"height": state["height"],
"width": state["width"],
"crop_region": state["crop_region"],
"seed": state["seed"],
"prompt_embeds": state["prompt_embeds"] if state["prompt_embeds"] is not None else fixed["default"],
"negative_embeds": fixed["negative"],
"prompt_ref_embeds": fixed["ref"],
"clip_ref": state["clip_ref"],
"clip_drive": state["clip_drive"],
"image_pixels": state["image_pixels"],
"num_segments": num_segments,
"seg_index": seg,
"cur_step": state["cur_step"],
"latents": state["latents"],
"sched_state": state["sched_state"],
"drive_latents": state["drive_latents"],
"prev_cond": state["prev_cond"],
"tail_frames": state["tail_frames"],
"timings": state["timings"],
"segment_pixels": segment_pixels,
}
job_path = os.path.join(workdir, "job.pt")
torch.save(job, job_path)
del job
header = (f"🔄 **در حال رندر با Wan-Animate-2 Distilled ({resolution}) روی کارت کامل — "
f"{num_segments} بخش، مهلت این مرحله تا {_format_seconds(window)} دقیقه**")
yield None, None, header
result, failure = yield from _run_gpu_with_recovery(
_min_window_seconds(state), run_gpu_window, lambda w: (resolution, job_path, w),
result_path, window, progress_path, header,
)
if result is not None:
_apply_result(state, result, workdir)
if state["seg_index"] < num_segments:
state_path = _save_state_file(state)
done_segments = state["seg_index"]
progress_line = (
f"پیشرفت: بخش **{done_segments}** از **{num_segments}** کامل شده؛ "
f"بخش جاری در گام **{state['cur_step']}** از **{NUM_INFERENCE_STEPS}** است.\n\n"
)
if failure is None:
status_msg = "⚠️ **این مرحله از رندر ذخیره شد.**\n\n" + progress_line + _CONTINUE_STEPS
else:
status_msg = (
"⚠️ **پنجره‌ی GPU این بار کامل انجام نشد، اما پیشرفت کار ذخیره شد.**\n\n"
+ progress_line
+ f"علت (از سمت ZeroGPU): `{failure[:400]}`\n\n"
+ _CONTINUE_STEPS
)
_round_sink.key = _progress_key(state)
yield None, state_path, status_msg
return
# همه‌ی بخش‌ها آماده‌اند: اتصال، افزودن صدا و خروجی نهایی
final_path = _finalize_output(state, workdir, "wan_animate2")
_round_sink.key = "final"
yield final_path, None, "✅ **رندر کل ویدیو با موفقیت به پایان رسید!**"
def _finalize_output(state, workdir, prefix):
seg_paths = [write_bytes(b, ".mp4", workdir) for b in state["segment_videos"]]
joined = os.path.join(workdir, "joined.mp4")
if len(seg_paths) == 1:
joined = seg_paths[0]
else:
concat_mp4s(seg_paths, joined)
final_path = os.path.join(OUTPUT_DIR, f"{prefix}_{uuid.uuid4().hex[:8]}.mp4")
if state.get("audio_bytes"):
audio_path = write_bytes(state["audio_bytes"], ".wav", workdir)
try:
combine_video_and_audio(joined, audio_path, final_path)
except Exception as e:
print(f"[audio] {e}", flush=True)
shutil.copyfile(joined, final_path)
else:
shutil.copyfile(joined, final_path)
return final_path
class RadioAnimated(gr.HTML):
def __init__(self, choices, value=None, **kwargs):
if not choices or len(choices) < 2:
raise ValueError("RadioAnimated requires at least 2 choices.")
if value is None:
value = choices[0]
uid = uuid.uuid4().hex[:8]
group_name = f"ra-{uid}"
inputs_html = "\n".join(
f"""
<input class="ra-input" type="radio" name="{group_name}" id="{group_name}-{i}" value="{c}">
<label class="ra-label" for="{group_name}-{i}">{c}</label>
"""
for i, c in enumerate(choices)
)
html_template = f"""
<div class="ra-wrap" data-ra="{uid}">
<div class="ra-inner">
<div class="ra-highlight"></div>
{inputs_html}
</div>
</div>
"""
js_on_load = r"""
(() => {
const wrap = element.querySelector('.ra-wrap');
const inner = element.querySelector('.ra-inner');
const highlight = element.querySelector('.ra-highlight');
const inputs = Array.from(element.querySelectorAll('.ra-input'));
const labels = Array.from(element.querySelectorAll('.ra-label'));
if (!inputs.length || !labels.length) return;
const choices = inputs.map(i => i.value);
const PAD = 6;
let currentIdx = 0;
function setHighlightByIndex(idx) {
currentIdx = idx;
const lbl = labels[idx];
if (!lbl) return;
const innerRect = inner.getBoundingClientRect();
const lblRect = lbl.getBoundingClientRect();
highlight.style.width = `${lblRect.width}px`;
const x = (lblRect.left - innerRect.left - PAD);
highlight.style.transform = `translateX(${x}px)`;
}
function setCheckedByValue(val, shouldTrigger=false) {
const idx = Math.max(0, choices.indexOf(val));
inputs.forEach((inp, i) => { inp.checked = (i === idx); });
requestAnimationFrame(() => setHighlightByIndex(idx));
props.value = choices[idx];
if (shouldTrigger) trigger('change', props.value);
}
setCheckedByValue(props.value ?? choices[0], false);
inputs.forEach((inp) => {
inp.addEventListener('change', () => setCheckedByValue(inp.value, true));
});
window.addEventListener('resize', () => setHighlightByIndex(currentIdx));
// وقتی ردیف مخفی دوباره نمایش داده می‌شود، جای هایلایت درست شود
if (window.ResizeObserver) {
new ResizeObserver(() => setHighlightByIndex(currentIdx)).observe(inner);
}
let last = props.value;
const syncFromProps = () => {
if (props.value !== last) {
last = props.value;
setCheckedByValue(last, false);
}
requestAnimationFrame(syncFromProps);
};
requestAnimationFrame(syncFromProps);
})();
"""
super().__init__(
value=value,
html_template=html_template,
js_on_load=js_on_load,
**kwargs
)
css = """
#col-container {
margin: 0 auto;
max-width: 1600px;
}
#step-column {
padding: 10px;
border-radius: 8px;
box-shadow: var(--card-shadow);
margin: 5px 10px 10px 10px;
}
.button-gradient {
background: linear-gradient(45deg, rgb(255, 65, 108), rgb(255, 75, 43), rgb(255, 155, 0), rgb(255, 65, 108)) 0% 0% / 400% 400%;
border: none;
padding: 14px 28px;
font-size: 16px;
font-weight: bold;
color: white;
border-radius: 10px;
cursor: pointer;
transition: 0.3s ease-in-out;
animation: 2s linear 0s infinite normal none running gradientAnimation;
box-shadow: rgba(255, 65, 108, 0.6) 0px 4px 10px;
}
.ra-wrap{
width: fit-content;
}
.ra-inner{
position: relative;
display: inline-flex;
align-items: center;
gap: 0;
padding: 6px;
background: #0b0b0b;
border-radius: 9999px;
overflow: hidden;
user-select: none;
}
.ra-input{
display: none;
}
.ra-label{
position: relative;
z-index: 2;
padding: 10px 18px;
font-family: ui-sans-serif, system-ui, -apple-system, Segoe UI, Roboto, Arial;
font-size: 14px;
font-weight: 600;
color: rgba(255,255,255,0.7);
cursor: pointer;
transition: color 180ms ease;
white-space: nowrap;
}
.ra-highlight{
position: absolute;
z-index: 1;
top: 6px;
left: 6px;
height: calc(100% - 12px);
border-radius: 9999px;
background: #8bff97;
transition: transform 200ms ease, width 200ms ease;
}
.ra-input:checked + .ra-label{
color: rgba(0,0,0,0.75);
}
#mode-row {
display: flex !important;
justify-content: center !important;
align-items: center !important;
width: 100% !important;
}
#mode-row > * {
flex: 0 0 auto !important;
width: auto !important;
min-width: 0 !important;
}
#mode-row .gr-html,
#mode-row .gradio-html,
#mode-row .prose,
#mode-row .block {
width: auto !important;
flex: 0 0 auto !important;
display: inline-block !important;
}
"""
MODEL_CHOICES = ["Motion Transfer", "Character Replace (Mix)"]
def _combine_resolution(model_value, motion_res, mix_res):
if model_value == MODEL_CHOICES[1]:
return f"mix-{normalize_mix_resolution(mix_res)}"
return normalize_resolution(motion_res)
def _existing(paths):
return [[p] for p in paths if os.path.exists(p)]
with gr.Blocks(title="Wan 2.2 Animate 2", delete_cache=(3600, 21600)) as demo:
with gr.Column(elem_id="col-container"):
with gr.Row():
gr.HTML(
f"""
<div style="text-align: center;">
<p style="font-size:16px; display: inline; margin: 0;">
<strong>Wan2.2-Animate-2-14B</strong>
</p>
<a href="https://huggingface.co/{MODEL_REPO}" style="display: inline-block; vertical-align: middle; margin-left: 0.5em;">
[Model]
</a>
</div>
"""
)
with gr.Row(elem_id="mode-row"):
model_choice = RadioAnimated(
choices=MODEL_CHOICES,
value=MODEL_CHOICES[0],
elem_id="radioanimated_model",
)
with gr.Row(elem_id="mode-row", visible=True) as motion_res_row:
resolution_choice = RadioAnimated(
choices=RESOLUTIONS,
value="720p",
elem_id="radioanimated_resolution",
)
with gr.Row(elem_id="mode-row", visible=False) as mix_res_row:
mix_resolution_choice = RadioAnimated(
choices=MIX_RESOLUTIONS,
value="720p",
elem_id="radioanimated_mix_resolution",
)
# مقداری که به API ارسال می‌شود: 720p / 1080p / mix-480p / mix-720p
job_resolution = gr.Textbox(value="720p", visible=False)
with gr.Row():
with gr.Column(elem_id="step-column"):
gr.HTML('<div><span style="font-size: 24px;">1. Upload a Video</span><br></div>')
input_video = gr.Video(label="Input Video", height=512)
video_examples = _existing(["./examples/martialart.mp4", "./examples/test_example.mp4", "./examples/dream.mp4"])
if video_examples:
gr.Examples(examples=video_examples, inputs=[input_video], cache_examples=False)
with gr.Column(elem_id="step-column"):
gr.HTML('<div><span style="font-size: 24px;">2. Upload a Ref Image</span><br></div>')
edited_frame = gr.Image(label="Ref Image", type="filepath", height=512)
with gr.Accordion("Prompt (اختیاری)", open=False):
prompt_text = gr.Textbox(
label="توضیح ظاهر شخصیت و پس‌زمینه‌ی تصویر مرجع",
placeholder="خالی بگذارید تا پرامپت پیش‌فرض مدل استفاده شود",
lines=3,
value="",
)
image_examples = _existing(["./examples/ali.png", "./examples/james.png", "./examples/amber.png", "./examples/ella.png"])
if image_examples:
gr.Examples(examples=image_examples, inputs=[edited_frame], cache_examples=False)
with gr.Column(elem_id="step-column"):
gr.HTML('<div><span style="font-size: 24px;">3. Wan Animate it!</span><br></div>')
output_video = gr.Video(label="Edited Video", height=512)
tracking_id_in = gr.File(
label="Upload state file (.pt) / آپلود فایل وضعیت (.pt) برای ادامه کار",
file_types=[".pt"],
type="filepath",
)
tracking_id_out = gr.File(
label="Download state file (.pt) / دانلود فایل وضعیت همگام‌سازی شده",
interactive=False,
)
transfer_state_btn = gr.Button("انتقال سریع فایل وضعیت به بخش ورودی 🔄", variant="secondary")
status_out = gr.Markdown(value="**وضعیت:** در انتظار شروع فرآیند...", label="System Status")
action_button = gr.Button("Wan Animate 🦆", variant="primary", elem_classes="button-gradient")
action_button.click(
fn=animate_scene,
inputs=[input_video, edited_frame, job_resolution, prompt_text, tracking_id_in],
outputs=[output_video, tracking_id_out, status_out],
)
def _on_model_change(model_value, motion_res, mix_res):
is_mix = model_value == MODEL_CHOICES[1]
return (
gr.update(visible=not is_mix),
gr.update(visible=is_mix),
_combine_resolution(model_value, motion_res, mix_res),
)
model_choice.change(
_on_model_change,
inputs=[model_choice, resolution_choice, mix_resolution_choice],
outputs=[motion_res_row, mix_res_row, job_resolution],
api_visibility="private",
)
for _res_component in (resolution_choice, mix_resolution_choice):
_res_component.change(
_combine_resolution,
inputs=[model_choice, resolution_choice, mix_resolution_choice],
outputs=[job_resolution],
api_visibility="private",
)
transfer_state_btn.click(
fn=do_transfer,
inputs=[tracking_id_out],
outputs=[tracking_id_in, status_out],
api_visibility="private",
)
if __name__ == "__main__":
demo.queue()
demo.launch(ssr_mode=False, mcp_server=True, css=css, theme=gr.themes.Ocean())