# ========================================================================== # 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""" """ for i, c in enumerate(choices) ) html_template = f"""
{inputs_html}
""" 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"""

Wan2.2-Animate-2-14B

[Model]
""" ) 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('
1. Upload a Video
') 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('
2. Upload a Ref Image
') 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('
3. Wan Animate it!
') 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())