Wan2.2-Animate2 / mix_engine.py
Transfer Bot
Moved to Hugging Face automatically
a8d7e70
Raw History Blame Contribute Delete
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
@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)