# ========================================================================== # موتور حالت «جایگزینی شخصیت» (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 @torch.no_grad() 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 @torch.inference_mode() 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]}" @torch.no_grad() 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)