Spaces:
Runtime error
Runtime error
Download mix_engine.py from Opera8/Wan2.2-Animate2: direct link, hf CLI and curl.
- Browser
- Download file 29.5 kB
-
https://huggingface.co/spaces/Opera8/Wan2.2-Animate2/resolve/main/mix_engine.py
- Command line
-
hf download hf://spaces/Opera8/Wan2.2-Animate2/mix_engine.py
-
curl -L -o mix_engine.py https://huggingface.co/spaces/Opera8/Wan2.2-Animate2/resolve/main/mix_engine.py
29.5 kB
| # ========================================================================== | |
| # موتور حالت «جایگزینی شخصیت» (Replacement / Mix) با مدل Wan2.2-Animate-14B | |
| # ========================================================================== | |
| # این ماژول عیناً مسیر رسمی مخزن Wan2.2 را برای حالت replace بازسازی میکند: | |
| # ۱. پیشپردازش: تشخیص فرد (YOLOv10m) + نقاط بدن/دست/صورت (ViTPose-H) با ONNX روی GPU، | |
| # ماسک شخصیت با SAM2 (بخشهای ۴۰۰ فریمی با نقاط کلیدی بدن)، گشاد کردن ماسک | |
| # (k=7, iterations=3, w_len=h_len=1)، برش صورت ۵۱۲×۵۱۲ و رسم اسکلت. | |
| # ۲. تولید: سگمنتهای ۷۷ فریمی با ۱ فریم همپوشان، ۲۰ گام UniPC (shift=5)، بدون CFG، | |
| # بههمراه LoRA نورپردازی (relighting) که روی وزنها ادغام میشود. | |
| # | |
| # نکات ZeroGPU: | |
| # - هیچکدام از import های onnxruntime / sam2 در سطح ماژول انجام نمیشود؛ فقط داخل پنجرهی GPU. | |
| # - ترنسفورمر این مدل در پروسهی اصلی بارگذاری نمیشود (نه RAM و نه سهم کارت مدل اصلی را | |
| # نمیگیرد)؛ در هر پنجره مستقیم از دیسک روی GPU خوانده و در پایان آزاد میشود. | |
| # ========================================================================== | |
| import os | |
| import re | |
| import sys | |
| import math | |
| import zlib | |
| import numpy as np | |
| import cv2 | |
| import torch | |
| import torch.nn.functional as F | |
| APP_DIR = os.path.dirname(os.path.abspath(__file__)) | |
| PREPROCESS_DIR = os.path.join(APP_DIR, "animate_preprocess") | |
| if PREPROCESS_DIR not in sys.path: | |
| sys.path.insert(0, PREPROCESS_DIR) | |
| os.environ.setdefault("MPLBACKEND", "Agg") | |
| MIX_MODEL_REPO = "Wan-AI/Wan2.2-Animate-14B-Diffusers" | |
| MIX_AUX_REPO = "Wan-AI/Wan2.2-Animate-14B" | |
| MIX_MODEL_PATTERNS = ["transformer/*", "image_encoder/*", "scheduler/*", "vae/*"] | |
| MIX_AUX_PATTERNS = [ | |
| "process_checkpoint/det/*", | |
| "process_checkpoint/pose2d/*", | |
| "process_checkpoint/sam2/*", | |
| "relighting_lora/*", | |
| ] | |
| MIX_DEFAULT_PROMPT = "视频中的人在做动作" | |
| MIX_SEGMENT_FRAMES = 77 # clip_len رسمی | |
| MIX_PREV_FRAMES = 1 # refert_num رسمی | |
| MIX_EFFECTIVE = MIX_SEGMENT_FRAMES - MIX_PREV_FRAMES | |
| MIX_FPS = 30 | |
| MIX_STEPS = 20 | |
| MIX_SAM_CHUNK = 400 | |
| MIX_FACE_SIZE = 512 | |
| MIX_MASK_K = 7 | |
| MIX_MASK_ITERATIONS = 3 | |
| MIX_LATENT_FRAMES = (MIX_SEGMENT_FRAMES - 1) // 4 + 1 # = 20 | |
| # کیفیت چهره | |
| MIX_REF_SRC_MAX_SIDE = 2048 # نسخهی باکیفیت تصویر مرجع که برای همترازی نگه داریم | |
| MIX_ALIGN_SECONDS = 25.0 # برآورد زمان همترازی مرجع + CLIP روی GPU | |
| MIX_FACE_KP = slice(23, 91) # ۶۸ نقطهی صورت در خروجی ۱۳۳ نقطهای ViTPose | |
| MIX_FACE_MIN_SCORE = 0.3 | |
| MIX_FACE_MIN_POINTS = 24 | |
| MIX_ALIGN_MAX_UPSCALE = 3.0 # حداکثر بزرگنمایی نسبت به نسخهی ذخیرهشدهی مرجع | |
| # شدت اثر حالت چهرهی شخص ویدیو (همان face_strength در WanVideoWrapper؛ رسمی = 1.0). | |
| # کمتر = چهره آرامتر و شبیهتر به عکس، ولی حرکت لب/حالت چهره کمرنگتر. بازهی منطقی 0.6 تا 1.0 | |
| MIX_FACE_STRENGTH = 0.8 | |
| # -------------------------------------------------------------------------- | |
| # هندسه | |
| # -------------------------------------------------------------------------- | |
| def mix_frame_size(video_w, video_h, area, divisor=16): | |
| """ | |
| اندازهی فریم مطابق resize_by_area رسمی. (تابع calculate_new_size در کد رسمی بهخاطر | |
| یک خطای برنامهنویسی همیشه استثنا میدهد و مسیر جایگزین زیر اجرا میشود.) | |
| """ | |
| aspect = video_w / video_h | |
| new_h = math.sqrt(area / aspect) | |
| new_w = area / new_h | |
| return int((new_h // divisor) * divisor), int((new_w // divisor) * divisor) | |
| def resize_frame(frame, height, width): | |
| from preprocess_utils import padding_resize | |
| h, w = frame.shape[:2] | |
| interp = cv2.INTER_AREA if (width * height < w * h) else cv2.INTER_LINEAR | |
| return padding_resize(frame, height=height, width=width, interpolation=interp) | |
| # -------------------------------------------------------------------------- | |
| # تصویر مرجع: کوچکسازی بدون aliasing + همترازی کادر با شخص داخل ویدیو | |
| # -------------------------------------------------------------------------- | |
| def reference_source(image_rgb): | |
| """نسخهی باکیفیت مرجع برای ذخیره در state (بزرگترین ضلع ≤ MIX_REF_SRC_MAX_SIDE، با INTER_AREA).""" | |
| img = np.ascontiguousarray(image_rgb[:, :, :3]).astype(np.uint8) | |
| h, w = img.shape[:2] | |
| scale = MIX_REF_SRC_MAX_SIDE / float(max(h, w)) | |
| if scale < 1.0: | |
| img = cv2.resize(img, (max(1, int(round(w * scale))), max(1, int(round(h * scale)))), | |
| interpolation=cv2.INTER_AREA) | |
| return img | |
| def reference_padded(image_rgb, height, width): | |
| """ | |
| همان padding_resize رسمی ولی با INTER_AREA در کوچکسازی. کد رسمی با INTER_LINEAR کوچک میکند که | |
| روی عکسهای بزرگ (مثلاً ۳۵۰۰ پیکسل) پوست و مو را دانهدانه و نویزی میکند. | |
| """ | |
| return resize_frame(np.ascontiguousarray(image_rgb[:, :, :3]).astype(np.uint8), height, width) | |
| def _face_stats(kps): | |
| """(cx, cy, size) از ۶۸ نقطهی صورت؛ size = ریشهی مساحت کادر نقاط. None اگر صورت معتبر نبود.""" | |
| pts = np.asarray(kps)[MIX_FACE_KP] | |
| ok = pts[:, 2] >= MIX_FACE_MIN_SCORE | |
| if int(ok.sum()) < MIX_FACE_MIN_POINTS: | |
| return None | |
| p = pts[ok, :2] | |
| if not np.all(np.isfinite(p)): | |
| return None | |
| x0, y0 = p.min(axis=0) | |
| x1, y1 = p.max(axis=0) | |
| fw, fh = float(x1 - x0), float(y1 - y0) | |
| if fw < 4 or fh < 4: | |
| return None | |
| return (float(x0 + x1) / 2.0, float(y0 + y1) / 2.0, math.sqrt(fw * fh)) | |
| def driver_face_target(kp2ds, height, width, max_samples=240): | |
| """میانهی مرکز و اندازهی صورت شخص داخل ویدیو (روی فریمهای پخش در کل ویدیو).""" | |
| n = len(kp2ds) | |
| if n == 0: | |
| return None | |
| idxs = np.unique(np.linspace(0, n - 1, num=min(n, max_samples)).round().astype(int)) | |
| stats = [] | |
| for i in idxs: | |
| st = _face_stats(kp2ds[i]) | |
| if st is not None and 0 <= st[0] < width and 0 <= st[1] < height: | |
| stats.append(st) | |
| if len(stats) < max(1, len(idxs) // 4): | |
| return None | |
| arr = np.asarray(stats) | |
| return float(np.median(arr[:, 0])), float(np.median(arr[:, 1])), float(np.median(arr[:, 2])) | |
| def align_reference(pose_models, ref_src, driver_kp2ds, height, width): | |
| """ | |
| تصویر مرجع را طوری مقیاس و جابهجا میکند که اندازه و جای صورتش با صورت شخص داخل ویدیو یکی شود | |
| (پسزمینهی خالی سیاه، مثل padding رسمی). اگر صورت در مرجع یا ویدیو پیدا نشود None برمیگرداند. | |
| خروجی: (uint8 H,W,3, info) | |
| """ | |
| target = driver_face_target(driver_kp2ds, height, width) | |
| if target is None: | |
| return None, "no driver face" | |
| ref_kp = estimate_kp2ds(pose_models, [ref_src])[0] | |
| ref = _face_stats(ref_kp) | |
| if ref is None: | |
| return None, "no reference face" | |
| rh, rw = ref_src.shape[:2] | |
| fit = min(height / rh, width / rw) # مقیاس padding_resize | |
| s = target[2] / ref[2] | |
| lo = fit * 0.2 | |
| hi = max(fit, min(MIX_ALIGN_MAX_UPSCALE, fit * 6.0)) | |
| s = float(min(max(s, lo), hi)) | |
| nw, nh = max(1, int(round(rw * s))), max(1, int(round(rh * s))) | |
| interp = cv2.INTER_AREA if s < 1.0 else cv2.INTER_CUBIC | |
| scaled = cv2.resize(ref_src, (nw, nh), interpolation=interp) | |
| ox = int(round(target[0] - s * ref[0])) | |
| oy = int(round(target[1] - s * ref[1])) | |
| x0, y0 = max(0, ox), max(0, oy) | |
| x1, y1 = min(width, ox + nw), min(height, oy + nh) | |
| if x1 - x0 < 16 or y1 - y0 < 16: | |
| return None, "placement outside frame" | |
| canvas = np.zeros((height, width, 3), np.uint8) | |
| canvas[y0:y1, x0:x1] = scaled[y0 - oy:y1 - oy, x0 - ox:x1 - ox] | |
| # اگر عکس مرجع (مثلاً پرتره تا سینه) بالاتر از پایین کادر تمام شود، لباس را با رنگ نرم ادامه بده | |
| # تا زیر بدنِ شخص ویدیو نوار سیاه نیفتد. | |
| if oy + nh == y1 and y1 < height: | |
| band = canvas[max(y0, y1 - 8):y1, x0:x1].astype(np.float32).mean(axis=0, keepdims=True) | |
| ext = np.repeat(band, height - y1, axis=0) | |
| sigma = max(2.0, (x1 - x0) / 60.0) | |
| ext = cv2.GaussianBlur(ext, (0, 0), sigmaX=sigma, sigmaY=sigma, borderType=cv2.BORDER_REPLICATE) | |
| canvas[y1:height, x0:x1] = np.clip(ext, 0, 255).astype(np.uint8) | |
| info = (f"scale={s:.3f} (fit={fit:.3f}) ref_face={ref[2]:.0f}px drv_face={target[2]:.0f}px " | |
| f"offset=({ox},{oy})") | |
| return canvas, info | |
| def clip_encode_on_gpu(model_dir, ref_u8_hw3, device): | |
| """ویژگی CLIP تصویر مرجعِ همتراز (همان clip_visual_encode پروسهی CPU) → bf16 روی CPU.""" | |
| from transformers import CLIPVisionModel | |
| from diffusers.modular_pipelines.wan_animate_2.encoders import clip_visual_encode | |
| encoder = CLIPVisionModel.from_pretrained(model_dir, subfolder="image_encoder", dtype=torch.float32) | |
| encoder = encoder.to(device).eval() | |
| try: | |
| x = torch.from_numpy(np.ascontiguousarray(ref_u8_hw3)).permute(2, 0, 1).to(device, torch.float32) | |
| x = x / 127.5 - 1.0 | |
| feats = clip_visual_encode(encoder, x, device, torch.float32) | |
| return feats.to("cpu", torch.bfloat16) | |
| finally: | |
| del encoder | |
| torch.cuda.empty_cache() | |
| def load_mix_vae(model_dir, device, fallback): | |
| """VAE خود مدل Animate-14B با float32 (مثل کد رسمی)؛ اگر نشد همان VAE اشتراکی.""" | |
| try: | |
| from diffusers import AutoencoderKLWan | |
| path = os.path.join(model_dir, "vae") | |
| if not os.path.isdir(path): | |
| raise FileNotFoundError(path) | |
| v = AutoencoderKLWan.from_pretrained(model_dir, subfolder="vae", torch_dtype=torch.float32) | |
| return v.to(device).eval(), True | |
| except Exception as e: | |
| print(f"[mix] float32 VAE unavailable ({type(e).__name__}: {e}); using shared VAE", flush=True) | |
| return fallback, False | |
| def mix_segment_plan(real_frame_len): | |
| """مطابق WanAnimatePipeline: تعداد سگمنتها با پدینگ بازتابی تا مضرب طول مؤثر.""" | |
| last = (real_frame_len - MIX_PREV_FRAMES) % MIX_EFFECTIVE | |
| padding = 0 if last == 0 else MIX_EFFECTIVE - last | |
| target = real_frame_len + padding | |
| return max(1, target // MIX_EFFECTIVE), target | |
| def reflect_index(p, n): | |
| """[0..n-1, n-2..0, 1..] (همان pad_video_frames در diffusers).""" | |
| if n <= 1: | |
| return 0 | |
| period = 2 * (n - 1) | |
| m = p % period | |
| return m if m < n else period - m | |
| def segment_frame_indices(seg, real_frame_len): | |
| start = seg * MIX_EFFECTIVE | |
| return [reflect_index(p, real_frame_len) for p in range(start, start + MIX_SEGMENT_FRAMES)] | |
| # -------------------------------------------------------------------------- | |
| # خواندن فریمها (ویدیوی ۳۰ فریم بر ثانیه) با کش محدود | |
| # -------------------------------------------------------------------------- | |
| class FrameSource: | |
| def __init__(self, video_path, height, width, real_frame_len, cache_frames=240): | |
| self.path = video_path | |
| self.height = height | |
| self.width = width | |
| self.real = real_frame_len | |
| self.cache_frames = cache_frames | |
| self.cap = None | |
| self.pos = 0 | |
| self.cache = {} | |
| self.last_valid = None | |
| def _open(self): | |
| if self.cap is not None: | |
| self.cap.release() | |
| self.cap = cv2.VideoCapture(self.path) | |
| self.pos = 0 | |
| def _read_next(self): | |
| ok, frame = self.cap.read() | |
| idx = self.pos | |
| self.pos += 1 | |
| if ok: | |
| frame = resize_frame(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB), self.height, self.width) | |
| self.last_valid = frame | |
| else: | |
| frame = self.last_valid if self.last_valid is not None else np.zeros( | |
| (self.height, self.width, 3), np.uint8) | |
| self.cache[idx] = frame | |
| keep = max(self.cache_frames, getattr(self, "_keep", 0)) | |
| for k in [k for k in self.cache if k < idx - keep]: | |
| del self.cache[k] | |
| return frame | |
| def get(self, indices): | |
| need = sorted(set(int(i) for i in indices)) | |
| self._keep = need[-1] - need[0] + 1 | |
| missing = [i for i in need if i not in self.cache] | |
| if missing: | |
| if self.cap is None or min(missing) < self.pos: | |
| self.cache.clear() | |
| self._open() | |
| target = max(missing) | |
| while self.pos <= target: | |
| if self.pos in need or self.pos >= min(missing): | |
| self._read_next() | |
| else: | |
| self.cap.grab() | |
| self.pos += 1 | |
| out = {i: self.cache[i] for i in need} | |
| return [out[int(i)] for i in indices] | |
| def close(self): | |
| if self.cap is not None: | |
| self.cap.release() | |
| self.cap = None | |
| self.cache.clear() | |
| # -------------------------------------------------------------------------- | |
| # نقاط کلیدی (ONNX روی GPU) | |
| # -------------------------------------------------------------------------- | |
| def load_pose_models(aux_dir): | |
| import onnxruntime | |
| try: | |
| if hasattr(onnxruntime, "preload_dlls"): | |
| onnxruntime.preload_dlls() | |
| except Exception as e: | |
| print(f"[mix] onnxruntime preload_dlls: {e}", flush=True) | |
| from pose2d import Yolo, ViTPose | |
| ckpt = os.path.join(aux_dir, "process_checkpoint") | |
| det = Yolo(os.path.join(ckpt, "det", "yolov10m.onnx"), "cuda") | |
| pose = ViTPose(os.path.join(ckpt, "pose2d", "vitpose_h_wholebody.onnx"), "cuda") | |
| print(f"[mix] onnx providers: det={det.session.get_providers()} pose={pose.session.get_providers()}", flush=True) | |
| return det, pose | |
| def estimate_kp2ds(models, frames_rgb): | |
| """معادل Pose2d.__call__ رسمی ولی خروجی خام [N, 133, 3] (مختصات پیکسلی).""" | |
| det, pose = models | |
| out = [] | |
| for frame in frames_rgb: | |
| image = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) # همان load_images رسمی | |
| img, shape = det.preprocess(image) | |
| bbox = det(img[None], shape[None])[0][0]["bbox"] | |
| img, center, scale = pose.preprocess(image, bbox) | |
| out.append(pose(img[None], center[None], scale[None])) | |
| return np.concatenate(out, 0).astype(np.float32) | |
| def metas_from_kp2ds(kp2ds, width, height): | |
| """load_pose_metas_from_kp2ds_seq رسمی (با محافظت برای فریم اول نامعتبر).""" | |
| from pose2d_utils import split_kp2ds_for_aa | |
| metas = [] | |
| last_body = None | |
| for kps in kp2ds: | |
| kps = kps.astype(np.float64).copy() | |
| kps[:, 0] /= width | |
| kps[:, 1] /= height | |
| body, lhand, rhand, face = split_kp2ds_for_aa(kps, ret_face=True) | |
| if body[:, :2].min(axis=1).max() < 0 and last_body is not None: | |
| body = last_body | |
| last_body = body | |
| metas.append({ | |
| "width": width, "height": height, | |
| "keypoints_body": body, "keypoints_left_hand": lhand, | |
| "keypoints_right_hand": rhand, "keypoints_face": face, | |
| }) | |
| return metas | |
| def face_bboxes(metas, height, width): | |
| from preprocess_utils import get_face_bboxes | |
| boxes = [] | |
| prev = None | |
| for meta in metas: | |
| box = None | |
| try: | |
| with np.errstate(all="ignore"): | |
| x1, x2, y1, y2 = get_face_bboxes(meta["keypoints_face"][:, :2], scale=1.3, | |
| image_shape=(height, width)) | |
| if x2 - x1 >= 2 and y2 - y1 >= 2: | |
| box = [x1, x2, y1, y2] | |
| except Exception: | |
| box = None | |
| if box is None: | |
| box = prev if prev is not None else [width // 4, 3 * width // 4, 0, height // 2] | |
| prev = box | |
| boxes.append(box) | |
| return np.asarray(boxes, dtype=np.int32) | |
| # -------------------------------------------------------------------------- | |
| # ماسک شخصیت با SAM2 (فقط داخل پنجرهی GPU) | |
| # -------------------------------------------------------------------------- | |
| def build_sam_predictor(aux_dir, device, config_name="sam2_hiera_l.yaml"): | |
| from hydra import compose | |
| from hydra.utils import instantiate | |
| from omegaconf import OmegaConf | |
| import sam2 # noqa: F401 (hydra config module) | |
| from sam2.build_sam import _load_checkpoint | |
| overrides = [ | |
| "++model._target_=sam2.sam2_video_predictor.SAM2VideoPredictor", | |
| "++model.sam_mask_decoder_extra_args.dynamic_multimask_via_stability=true", | |
| "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_delta=0.05", | |
| "++model.sam_mask_decoder_extra_args.dynamic_multimask_stability_thresh=0.98", | |
| "++model.binarize_mask_from_pts_for_mem_enc=true", | |
| "++model.fill_hole_area=8", | |
| ] | |
| cfg = compose(config_name=config_name, overrides=overrides) | |
| OmegaConf.resolve(cfg) | |
| model = instantiate(cfg.model, _recursive_=True) | |
| ckpt = os.path.join(aux_dir, "process_checkpoint", "sam2", "sam2_hiera_large.pt") if aux_dir else None | |
| if ckpt and os.path.exists(ckpt): | |
| _load_checkpoint(model, ckpt) | |
| return model.to(device).eval() | |
| def _sam_init_state(predictor, frames, device): | |
| """init_state_v2 رسمی Wan (فریمها از حافظه) با ساختار نسخهی فعلی SAM2.""" | |
| from collections import OrderedDict | |
| size = predictor.image_size | |
| mean = torch.tensor((0.485, 0.456, 0.406), dtype=torch.float32)[:, None, None] | |
| std = torch.tensor((0.229, 0.224, 0.225), dtype=torch.float32)[:, None, None] | |
| images = torch.zeros(len(frames), 3, size, size, dtype=torch.float32) | |
| from PIL import Image | |
| for n, frame in enumerate(frames): | |
| pil = Image.fromarray(frame.astype(np.uint8)) | |
| arr = np.array(pil.convert("RGB").resize((size, size))) / 255.0 | |
| images[n] = torch.from_numpy(arr).permute(2, 0, 1) | |
| video_height, video_width = frames[0].shape[:2] | |
| images = ((images - mean) / std).to(device) | |
| state = { | |
| "images": images, | |
| "num_frames": len(frames), | |
| "offload_video_to_cpu": False, | |
| "offload_state_to_cpu": False, | |
| "video_height": video_height, | |
| "video_width": video_width, | |
| "device": device, | |
| "storage_device": device, | |
| "point_inputs_per_obj": {}, | |
| "mask_inputs_per_obj": {}, | |
| "cached_features": {}, | |
| "constants": {}, | |
| "obj_id_to_idx": OrderedDict(), | |
| "obj_idx_to_id": OrderedDict(), | |
| "obj_ids": [], | |
| "output_dict_per_obj": {}, | |
| "temp_output_dict_per_obj": {}, | |
| "frames_tracked_per_obj": {}, | |
| } | |
| predictor._get_image_feature(state, frame_idx=0, batch_size=1) | |
| return state | |
| def sam_chunk_masks(predictor, frames, metas, device): | |
| """یک بخش از get_mask رسمی (ProcessPipeline.get_mask) → لیست ماسک uint8.""" | |
| import warnings | |
| warnings.filterwarnings("ignore", message=".*post-processing step.*") | |
| n = len(frames) | |
| if n == 0: | |
| return [] | |
| key_frame_num = 4 if n > 4 else 1 | |
| step = max(1, len(metas) // key_frame_num) | |
| key_indices = list(range(0, len(metas), step)) | |
| key_points_index = [0, 1, 2, 5, 8, 11, 10, 13] | |
| wh = np.array([[metas[0]["width"], metas[0]["height"]]]) | |
| state = _sam_init_state(predictor, frames, device) | |
| predictor.reset_state(state) | |
| for idx in key_indices: | |
| body = metas[idx]["keypoints_body"] | |
| pts = np.array([body[i] for i in key_points_index if body[i] is not None])[:, :2] | |
| points = (pts * wh).astype(np.int32) | |
| labels = np.array([1] * points.shape[0], np.int32) | |
| predictor.add_new_points(inference_state=state, frame_idx=idx, obj_id=1, | |
| points=points, labels=labels) | |
| segments = {} | |
| for out_idx, out_ids, out_logits in predictor.propagate_in_video(state): | |
| segments[out_idx] = (out_logits[0] > 0.0).cpu().numpy()[0].astype(np.uint8) | |
| blank = np.zeros(frames[0].shape[:2], np.uint8) | |
| masks = [segments.get(i, blank) for i in range(n)] | |
| del state | |
| return masks | |
| def augment_mask(mask): | |
| """get_mask_body_img(k=7, iterations=3) + get_aug_mask(w_len=1, h_len=1) رسمی.""" | |
| kernel = np.ones((MIX_MASK_K, MIX_MASK_K), np.uint8) | |
| dil = cv2.dilate(mask.astype(np.uint8), kernel, iterations=MIX_MASK_ITERATIONS) | |
| ys, xs = np.nonzero(dil) | |
| if len(xs) == 0: | |
| return dil | |
| x0, x1, y0, y1 = xs.min(), xs.max(), ys.min(), ys.max() | |
| if x1 > x0 and y1 > y0 and dil[y0:y1, x0:x1].sum() > 0: | |
| dil[y0:y1, x0:x1] = 1 | |
| return dil | |
| def pack_masks(masks): | |
| arr = np.stack(masks).astype(bool) | |
| return zlib.compress(np.packbits(arr).tobytes(), 6) | |
| def unpack_masks(blob, count, height, width): | |
| bits = np.frombuffer(zlib.decompress(blob), dtype=np.uint8) | |
| return np.unpackbits(bits)[: count * height * width].reshape(count, height, width) | |
| # -------------------------------------------------------------------------- | |
| # ورودیهای پیکسلی یک سگمنت (روی CPU، uint8) | |
| # -------------------------------------------------------------------------- | |
| def segment_pixels(frames_src, seg, real_frame_len, height, width, kp2ds, bboxes, mask_blobs): | |
| from pose2d_utils import AAPoseMeta | |
| from human_visualization import draw_aapose_by_meta_new | |
| idxs = segment_frame_indices(seg, real_frame_len) | |
| frames = frames_src.get(idxs) | |
| metas_all = metas_from_kp2ds(kp2ds, width, height) | |
| chunk_cache = {} | |
| def mask_of(i): | |
| c = i // MIX_SAM_CHUNK | |
| if c not in chunk_cache: | |
| count = min(MIX_SAM_CHUNK, real_frame_len - c * MIX_SAM_CHUNK) | |
| chunk_cache[c] = unpack_masks(mask_blobs[c], count, height, width) | |
| return chunk_cache[c][i - c * MIX_SAM_CHUNK] | |
| pose, face, bg, mask = [], [], [], [] | |
| for frame, i in zip(frames, idxs): | |
| canvas = np.zeros((height, width, 3), np.uint8) | |
| pose.append(draw_aapose_by_meta_new(canvas, AAPoseMeta.from_humanapi_meta(metas_all[i]))) | |
| x1, x2, y1, y2 = [int(v) for v in bboxes[i]] | |
| crop = frame[y1:y2, x1:x2] | |
| if crop.size == 0: | |
| crop = frame | |
| face.append(cv2.resize(crop, (MIX_FACE_SIZE, MIX_FACE_SIZE))) | |
| m = mask_of(i) | |
| bg.append(frame * (1 - m[:, :, None])) | |
| mask.append(m) | |
| def to_t(lst): | |
| return torch.from_numpy(np.ascontiguousarray(np.stack(lst))).permute(3, 0, 1, 2).unsqueeze(0) # 1,C,T,H,W | |
| return { | |
| "pose": to_t(pose), | |
| "face": to_t(face), | |
| "bg": to_t(bg), | |
| "mask": torch.from_numpy(np.stack(mask)).unsqueeze(0).unsqueeze(0), # 1,1,T,H,W (0/1) | |
| } | |
| def u8_to_pm1(x, device, dtype=torch.float32): | |
| return x.to(device).to(dtype) / 127.5 - 1.0 | |
| # -------------------------------------------------------------------------- | |
| # ترنسفورمر + LoRA نورپردازی | |
| # -------------------------------------------------------------------------- | |
| _LORA_ATTN_MAP = {"q": "to_q", "k": "to_k", "v": "to_v", "o": "to_out.0", "k_img": "add_k_proj", "v_img": "add_v_proj"} | |
| _LORA_RE = re.compile( | |
| r"(?:^|\.)blocks\.(\d+)\.(self_attn|cross_attn|ffn)\.(q|k|v|o|k_img|v_img|0|2)\." | |
| r"(lora_A|lora_B|lora_down|lora_up)(?:\.default)?\.weight$" | |
| ) | |
| def _lora_target_name(block, part, sub): | |
| if part == "ffn": | |
| return f"blocks.{block}.ffn.net.0.proj" if sub == "0" else f"blocks.{block}.ffn.net.2" | |
| attn = "attn1" if part == "self_attn" else "attn2" | |
| if sub not in _LORA_ATTN_MAP: | |
| return None | |
| return f"blocks.{block}.{attn}.{_LORA_ATTN_MAP[sub]}" | |
| def merge_relighting_lora(transformer, aux_dir, device, alpha=128.0): | |
| from safetensors.torch import load_file | |
| path = os.path.join(aux_dir, "relighting_lora", "adapter_model.safetensors") | |
| if not os.path.exists(path): | |
| print("[mix] relighting LoRA پیدا نشد؛ بدون آن ادامه میدهیم", flush=True) | |
| return 0 | |
| sd = load_file(path, device="cpu") | |
| pairs = {} | |
| skipped = 0 | |
| for key, tensor in sd.items(): | |
| m = _LORA_RE.search(key) | |
| if not m: | |
| skipped += 1 | |
| continue | |
| block, part, sub, kind = m.groups() | |
| target = _lora_target_name(block, part, sub) | |
| if target is None: | |
| skipped += 1 | |
| continue | |
| pairs.setdefault(target, {})["A" if kind in ("lora_A", "lora_down") else "B"] = tensor | |
| modules = dict(transformer.named_modules()) | |
| merged = 0 | |
| for name, ab in pairs.items(): | |
| mod = modules.get(name) | |
| if mod is None or "A" not in ab or "B" not in ab: | |
| skipped += 1 | |
| continue | |
| a = ab["A"].to(device, torch.float32) | |
| b = ab["B"].to(device, torch.float32) | |
| rank = a.shape[0] | |
| delta = (b @ a) * (alpha / rank) | |
| if delta.shape != mod.weight.shape: | |
| skipped += 1 | |
| continue | |
| mod.weight.add_(delta.to(mod.weight.dtype)) | |
| merged += 1 | |
| del a, b, delta | |
| del sd | |
| print(f"[mix] relighting LoRA merged into {merged} layers (skipped {skipped})", flush=True) | |
| return merged | |
| def apply_face_strength(transformer, strength=MIX_FACE_STRENGTH): | |
| """ | |
| خروجی بلوکهای face_adapter (که حالت چهرهی شخص ویدیو را تزریق میکنند) را در strength ضرب میکند. | |
| در diffusers این خروجی مستقیم با hidden_states جمع میشود، پس این دقیقاً معادل | |
| hidden_states + strength * face_adapter(...) است. | |
| """ | |
| adapters = getattr(transformer, "face_adapter", None) | |
| if adapters is None or abs(float(strength) - 1.0) < 1e-6: | |
| return 0 | |
| s = float(strength) | |
| def _scale(_module, _inputs, output): | |
| return output * s | |
| for block in adapters: | |
| block.register_forward_hook(_scale) | |
| print(f"[mix] face strength {s:.2f} on {len(adapters)} face adapter blocks", flush=True) | |
| return len(adapters) | |
| def load_mix_transformer(model_dir, aux_dir, device): | |
| from diffusers import WanAnimateTransformer3DModel | |
| transformer = WanAnimateTransformer3DModel.from_pretrained( | |
| model_dir, subfolder="transformer", torch_dtype=torch.bfloat16, device_map=str(device), | |
| ) | |
| transformer.eval().requires_grad_(False) | |
| merge_relighting_lora(transformer, aux_dir, device) | |
| apply_face_strength(transformer) | |
| return transformer | |
| # -------------------------------------------------------------------------- | |
| # شرطهای سگمنت (مطابق animate.py رسمی در حالت replace) | |
| # -------------------------------------------------------------------------- | |
| def i2v_mask(lat_t, lat_h, lat_w, mask_len=1, mask_pixels=None, device="cuda"): | |
| """mask_pixels: [T_pixel, lat_h, lat_w] یا None. خروجی [4, lat_t, lat_h, lat_w].""" | |
| if mask_pixels is not None: | |
| msk = mask_pixels.clone().to(device=device, dtype=torch.float32).unsqueeze(0) | |
| else: | |
| msk = torch.zeros(1, (lat_t - 1) * 4 + 1, lat_h, lat_w, device=device) | |
| msk[:, :mask_len] = 1 | |
| msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) | |
| msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) | |
| return msk.transpose(1, 2)[0] | |
| def reference_condition(encode_fn, ref_pixels_u8, device): | |
| """y_ref: ماسک ۱ + لیتنت تصویر مرجع → [20, 1, h, w].""" | |
| ref = u8_to_pm1(ref_pixels_u8, device).unsqueeze(0).unsqueeze(2) # 1,3,1,H,W | |
| lat = encode_fn(ref)[0] | |
| msk = i2v_mask(1, lat.shape[2], lat.shape[3], 1, device=device) | |
| return torch.cat([msk, lat], dim=0) | |
| def segment_conditions(encode_fn, pixels, tail_frames, is_first, height, width, device): | |
| """pose_latents [1,16,20,h,w] و y_reft [20,20,h,w].""" | |
| lat_h, lat_w = height // 8, width // 8 | |
| pose = encode_fn(u8_to_pm1(pixels["pose"], device)) | |
| bg = u8_to_pm1(pixels["bg"], device) | |
| if not is_first and tail_frames is not None: | |
| prev = tail_frames.to(device, torch.float32).unsqueeze(0) # 1,3,1,H,W | |
| cond_video = torch.cat([prev[:, :, :MIX_PREV_FRAMES], bg[:, :, MIX_PREV_FRAMES:]], dim=2) | |
| mask_len = MIX_PREV_FRAMES | |
| else: | |
| cond_video = bg | |
| mask_len = 0 | |
| y_lat = encode_fn(cond_video)[0] | |
| del bg, cond_video | |
| m = 1.0 - pixels["mask"][0, 0].to(device, torch.float32) # T,H,W | |
| m = F.interpolate(m.unsqueeze(1), size=(lat_h, lat_w), mode="nearest")[:, 0] | |
| msk = i2v_mask(MIX_LATENT_FRAMES, lat_h, lat_w, mask_len, mask_pixels=m, device=device) | |
| return pose, torch.cat([msk, y_lat], dim=0) | |