Wan2.2-Animate2 / animate2_engine.py
Transfer Bot
Moved to Hugging Face automatically
a8d7e70
Raw History Blame Contribute Delete
20.9 kB
# ==========================================================================
# موتور استنتاج Wan-Animate-2 (Wan2.2-Animate-2-14B) برای ZeroGPU
# ==========================================================================
# این ماژول هسته‌ی مدل جدید را با کامپوننت‌های رسمی diffusers اجرا می‌کند، اما
# حلقه‌ی دینویز را خودمان می‌نویسیم تا بتوان وسط هر سگمنت (بعد از هر گام)
# متوقف شد، وضعیت را ذخیره کرد و در پنجره‌ی GPU بعدی دقیقاً از همان‌جا ادامه داد.
#
# دو تغییر نسبت به مسیر پیش‌فرض diffusers (هر دو از نظر ریاضی معادل‌اند):
# ۱. توجه «in-context» حالت cached در diffusers فقط با flex-attention کامپایل‌شده
# کار می‌کند (torch.compile روی هر فورک ZeroGPU چند ده ثانیه هزینه دارد).
# اینجا همان ماسک را بدون ماسک و بدون کامپایل محاسبه می‌کنیم: یک فراخوانی
# flash روی کل توکن‌های تولید + یک فراخوانی flash (batch روی فریم‌ها) روی
# توکن‌های مرجع همان فریم، و ترکیب دقیق دو softmax با logsumexp.
# ۲. کش مرجع به‌صورت تطبیقی: اگر حافظه‌ی GPU جا داشته باشد K/V آماده (سریع‌ترین حالت)
# نگه داشته می‌شود؛ وگرنه فقط ورودی لایه (نصف حجم K+V) ذخیره و K/V در هر گام دقیقاً
# از نو ساخته می‌شود؛ و اگر باز هم جا نبود، لایه‌های باقی روی CPU می‌روند.
# (سخت‌افزار فعلی ZeroGPU کارت ۹۶ گیگابایتی است؛ کش K/V در ۱۰۸۰p حدود ۱۴۰ گیگابایت است.)
# ==========================================================================
import math
import numpy as np
import torch
import torch.nn.functional as F
from diffusers.models.transformers.transformer_wan_animate_2 import (
WanAnimate2AttnProcessor,
_get_qkv_projections,
)
from diffusers.modular_pipelines.wan_animate_2.encoders import get_i2v_mask
# پرامپت‌های پیش‌فرض (برگرفته از مخزن رسمی Wan-Animate-2 و diffusers)
DEFAULT_PROMPT = "视频中的人在做动作, 背景静止"
DEFAULT_PROMPT_REF = "人物动作的参考视频"
DEFAULT_NEGATIVE_PROMPT = (
"过曝,静态,细节模糊不清,字幕,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,"
"多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,三条腿"
)
SEGMENT_FRAME_LENGTH = 81 # طول هر سگمنت (مطابق پیش‌فرض رسمی)
PREV_SEGMENT_COND_FRAMES = 1 # فریم‌های هم‌پوشان بین سگمنت‌ها
EFFECTIVE_SEGMENT = SEGMENT_FRAME_LENGTH - PREV_SEGMENT_COND_FRAMES
OUTPUT_FPS = 24
ROPE_CHUNK_TOKENS = 16384
# --------------------------------------------------------------------------
# RoPE تکه‌تکه (خروجی دقیقاً برابر rope_apply اصلی، با مصرف حافظه‌ی محدود)
# --------------------------------------------------------------------------
def rope_apply_chunked(x, grid_sizes, freqs, time_stride=1, out_dtype=torch.float32, chunk=ROPE_CHUNK_TOKENS):
n, c = x.size(2), x.size(3) // 2
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
freqs_i = torch.cat(
[
freqs[0][: f * time_stride : time_stride].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),
],
dim=-1,
).reshape(seq_len, 1, -1)
out_i = torch.empty(x.size(1), n, 2 * c, dtype=out_dtype, device=x.device)
for s in range(0, seq_len, chunk):
e = min(s + chunk, seq_len)
xc = torch.view_as_complex(x[i, s:e].to(torch.float64).reshape(e - s, n, -1, 2))
out_i[s:e] = torch.view_as_real(xc * freqs_i[s:e]).flatten(2).float().to(out_dtype)
del xc
if seq_len < x.size(1):
out_i[seq_len:] = x[i, seq_len:].float().to(out_dtype)
output.append(out_i)
return torch.stack(output)
# --------------------------------------------------------------------------
# توجه با خروجی logsumexp — ورودی/خروجی با چیدمان [B, L, H, D]
# --------------------------------------------------------------------------
def _attn_with_lse(q, k, v):
qt, kt, vt = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
lq = qt.shape[2]
if qt.is_cuda:
try:
res = torch.ops.aten._scaled_dot_product_flash_attention(qt, kt, vt, 0.0, False, False)
return res[0].transpose(1, 2), res[1][..., :lq].transpose(1, 2)
except Exception:
pass
try:
res = torch.ops.aten._scaled_dot_product_efficient_attention(qt, kt, vt, None, True, 0.0, False)
return res[0].transpose(1, 2), res[1][..., :lq].transpose(1, 2)
except Exception:
pass
# مسیر مرجع (CPU / تست): محاسبه‌ی صریح
scale = 1.0 / math.sqrt(qt.shape[-1])
scores = torch.matmul(qt.float(), kt.float().transpose(-1, -2)) * scale
lse = torch.logsumexp(scores, dim=-1)
out = torch.matmul(torch.softmax(scores, dim=-1), vt.float()).to(q.dtype)
return out.transpose(1, 2), lse.transpose(1, 2)
# --------------------------------------------------------------------------
# پردازشگر توجه جایگزین (بدون flex و بدون torch.compile)
# --------------------------------------------------------------------------
class Animate2ExactAttnProcessor(WanAnimate2AttnProcessor):
"""
extract: توجه کامل روی توکن‌های مرجع + ذخیره‌ی K (از قبل RoPE شده) و V در کش.
(کلید RoPE‌شده در extract دقیقاً همان مقداری است که حالت cached در
هر گام از نو محاسبه می‌کرد؛ پس ذخیره‌ی مستقیم آن معادل است.)
cached : هر توکن فریم f از بخش تولید، به همه‌ی توکن‌های تولید + توکن‌های مرجع
فریم f-1 توجه می‌کند (فریم ۰ = جای تصویر مرجع، فقط تولید).
"""
def __init__(self, offload_policy):
self.offload_policy = offload_policy
def __call__(
self,
attn,
hidden_states,
rotary_emb,
grid_sizes,
kv_cache,
kv_cache_mode,
rope_stride=1,
reference_rotary_emb=None,
reference_grid_sizes=None,
reference_rope_stride=1,
attention_mask=None,
origin_latent_frames=None,
origin_latent_hw=None,
):
query, key, value = _get_qkv_projections(attn, hidden_states, None)
query = attn.norm_q(query)
key = attn.norm_k(key)
query = query.unflatten(2, (attn.heads, -1))
key = key.unflatten(2, (attn.heads, -1))
value = value.unflatten(2, (attn.heads, -1))
if kv_cache_mode == "extract":
kv_bytes = key.numel() * value.element_size() + value.nbytes
mode, device = self.offload_policy.place(kv_bytes, hidden_states.nbytes, query.device)
if mode == "x":
# فقط ورودی لایه ذخیره می‌شود؛ K/V در حالت cached دقیقاً از همین ساخته می‌شوند
kv_cache.store(hidden_states.detach().to(device), None)
query = rope_apply_chunked(query, grid_sizes, rotary_emb, rope_stride, out_dtype=value.dtype)
key = rope_apply_chunked(key, grid_sizes, rotary_emb, rope_stride, out_dtype=value.dtype)
hs = F.scaled_dot_product_attention(
query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
).transpose(1, 2)
if mode == "kv":
if device.type == query.device.type:
kv_cache.store(key, value.contiguous())
else:
kv_cache.store(key.to(device), value.to(device))
hidden_states = hs
elif kv_cache_mode == "cached":
query = rope_apply_chunked(query, grid_sizes, rotary_emb, rope_stride, out_dtype=value.dtype)
key = rope_apply_chunked(key, grid_sizes, rotary_emb, rope_stride, out_dtype=value.dtype)
key_ref, value_ref = kv_cache.get()
if value_ref is None:
x_ref = key_ref.to(query.device) if key_ref.device != query.device else key_ref
key_ref, value_ref = _reference_kv_from_input(
attn, x_ref, reference_grid_sizes, reference_rotary_emb, reference_rope_stride, value.dtype
)
del x_ref
elif key_ref.device != query.device:
key_ref = key_ref.to(query.device)
value_ref = value_ref.to(query.device)
frames, height, width = grid_sizes[0].tolist()
ref_frames, ref_h, ref_w = reference_grid_sizes[0].tolist()
hw, ref_hw = height * width, ref_h * ref_w
if hw != ref_hw or ref_frames != frames - 1:
raise RuntimeError("ناهمخوانی شبکه‌ی توکن‌های تولید و مرجع")
bsz, total, heads, head_dim = query.shape
valid = frames * hw
q = query[:, :valid]
out_all, lse_all = _attn_with_lse(q, key[:, :valid], value[:, :valid])
q_f = q[:, hw:].reshape(ref_frames, hw, heads, head_dim)
k_f = key_ref[:, : ref_frames * hw].reshape(ref_frames, hw, heads, head_dim)
v_f = value_ref[:, : ref_frames * hw].reshape(ref_frames, hw, heads, head_dim)
out_ref, lse_ref = _attn_with_lse(q_f, k_f, v_f)
del key_ref, value_ref, k_f, v_f
out_all = out_all.reshape(frames, hw, heads, head_dim)
lse_all = lse_all.reshape(frames, hw, heads)
for f in range(ref_frames):
w_gen = torch.sigmoid(lse_all[f + 1].float() - lse_ref[f].float()).unsqueeze(-1)
mixed = out_all[f + 1].float() * w_gen + out_ref[f].float() * (1.0 - w_gen)
out_all[f + 1].copy_(mixed.to(out_all.dtype))
del w_gen, mixed
del out_ref, lse_ref, lse_all
hidden_states = out_all
hidden_states = hidden_states.reshape(bsz, valid, heads, head_dim)
if valid < total:
hidden_states = torch.cat([hidden_states, query[:, valid:]], dim=1)
else:
raise ValueError(f"kv_cache_mode نامعتبر: {kv_cache_mode}")
hidden_states = hidden_states.flatten(2, 3).type_as(query)
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
return hidden_states
def _reference_kv_from_input(attn, x_ref, grid_sizes, rotary_emb, rope_stride, out_dtype):
"""K/V مرجع از ورودی ذخیره‌شده‌ی لایه — دقیقاً همان محاسبه‌ی مسیر extract."""
if attn.fused_projections:
_, key, value = attn.to_qkv(x_ref).chunk(3, dim=-1)
else:
key = attn.to_k(x_ref)
value = attn.to_v(x_ref)
key = attn.norm_k(key).unflatten(2, (attn.heads, -1))
value = value.unflatten(2, (attn.heads, -1))
key = rope_apply_chunked(key, grid_sizes, rotary_emb, rope_stride, out_dtype=out_dtype)
return key, value
class KVOffloadPolicy:
"""
برای هر لایه تصمیم می‌گیرد کش مرجع چه شکلی باشد و کجا بماند:
("kv", gpu) : K/V آماده روی GPU — سریع‌ترین
("x", gpu) : فقط ورودی لایه روی GPU — نصف حجم، K/V در هر گام ساخته می‌شود
("x", cpu) : ورودی لایه روی CPU — وقتی GPU دیگر جا ندارد
برنامه در اولین لایه‌ی پاس extract از روی حافظه‌ی آزاد واقعی ساخته می‌شود و قبل از
هر ذخیره‌سازی روی GPU دوباره چک می‌شود.
"""
def __init__(self):
self.force = None # فقط برای تست: "kv_gpu" | "x_gpu" | "x_cpu"
self.reset(0, 1)
def reset(self, reserve_bytes, num_layers):
self.reserve_bytes = int(reserve_bytes)
self.num_layers = int(num_layers)
self.layer = 0
self.plan = None
self.gpu_layers = 0
self.cpu_layers = 0
self.x_layers = 0
@staticmethod
def _free_bytes():
free, _ = torch.cuda.mem_get_info()
return free + torch.cuda.memory_reserved() - torch.cuda.memory_allocated()
def _make_plan(self, kv_bytes, x_bytes):
n = self.num_layers
budget = self._free_bytes() - self.reserve_bytes
if n * kv_bytes <= budget:
n_kv, n_x_gpu = n, 0
elif n * x_bytes <= budget:
n_kv = int((budget - n * x_bytes) // max(1, kv_bytes - x_bytes))
n_kv = max(0, min(n, n_kv))
n_x_gpu = n - n_kv
else:
n_kv = 0
n_x_gpu = max(0, min(n, int(budget // max(1, x_bytes))))
return n_kv, n_x_gpu
def place(self, kv_bytes, x_bytes, device):
idx = self.layer
self.layer += 1
cpu = torch.device("cpu")
if device.type != "cuda":
mode = {"kv_gpu": "kv", "x_gpu": "x", "x_cpu": "x"}.get(self.force, "kv")
if mode == "x":
self.x_layers += 1
return mode, device
if self.force:
mode = "kv" if self.force == "kv_gpu" else "x"
target = cpu if self.force == "x_cpu" else device
else:
if self.plan is None:
self.plan = self._make_plan(kv_bytes, x_bytes)
n_kv, n_x_gpu = self.plan
if idx < n_kv:
mode, target = "kv", device
elif idx < n_kv + n_x_gpu:
mode, target = "x", device
else:
mode, target = "x", cpu
if target.type == "cuda":
need = kv_bytes if mode == "kv" else x_bytes
if self._free_bytes() - need <= self.reserve_bytes:
if mode == "kv" and self._free_bytes() - x_bytes > self.reserve_bytes:
mode = "x"
else:
mode, target = "x", cpu
if mode == "x":
self.x_layers += 1
if target.type == "cuda":
self.gpu_layers += 1
else:
self.cpu_layers += 1
return mode, target
TEST_MODES = ["kv_gpu", "x_gpu", "x_cpu"]
def install_exact_attention(transformer):
policy = KVOffloadPolicy()
for block in transformer.blocks:
block.self_attn.set_processor(Animate2ExactAttnProcessor(policy))
# ماسک flex دیگر لازم نیست (و ساختنش torch.compile را صدا می‌زند)
transformer.create_mask = lambda *args, **kwargs: None
return policy
# --------------------------------------------------------------------------
# متن: T5 بدون پدینگ (خروجی توکن‌های واقعی با نسخه‌ی پدشده یکسان است)
# --------------------------------------------------------------------------
@torch.no_grad()
def encode_prompt_t5(text_encoder, tokenizer, prompt, max_sequence_length=512):
inputs = tokenizer(
[prompt],
padding=False,
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
)
ids = inputs.input_ids.to(text_encoder.device)
mask = inputs.attention_mask.to(text_encoder.device)
hidden = text_encoder(ids, mask).last_hidden_state[0]
out = hidden.new_zeros(max_sequence_length, hidden.size(-1))
out[: hidden.size(0)] = hidden
return out.to(torch.bfloat16).cpu() # [512, 4096]
# --------------------------------------------------------------------------
# وضعیت زمان‌بند (برای ادامه‌ی دقیق وسط سگمنت)
# --------------------------------------------------------------------------
_SCHED_KEYS = [
"model_outputs", "lower_order_nums", "_step_index", "_begin_index",
"last_sample", "this_order", "timestep_list", "order_list",
]
def _map_tensors(obj, fn):
if isinstance(obj, torch.Tensor):
return fn(obj)
if isinstance(obj, np.ndarray):
return fn(torch.from_numpy(np.ascontiguousarray(obj)))
if isinstance(obj, np.generic):
return obj.item()
if isinstance(obj, list):
return [_map_tensors(o, fn) for o in obj]
if isinstance(obj, tuple):
return tuple(_map_tensors(o, fn) for o in obj)
return obj
def scheduler_state_dict(scheduler):
state = {}
for k in _SCHED_KEYS:
if k in scheduler.__dict__:
state[k] = _map_tensors(scheduler.__dict__[k], lambda t: t.detach().to("cpu"))
return state
def load_scheduler_state(scheduler, state, device):
for k, v in (state or {}).items():
setattr(scheduler, k, _map_tensors(v, lambda t: t.to(device)))
# --------------------------------------------------------------------------
# محاسبات هندسه‌ی سگمنت‌ها (مطابق پیش‌پردازش diffusers + پدینگ زیگزاگ)
# --------------------------------------------------------------------------
def segment_plan(real_frame_len):
if real_frame_len > PREV_SEGMENT_COND_FRAMES:
leftover = (real_frame_len - PREV_SEGMENT_COND_FRAMES) % EFFECTIVE_SEGMENT
else:
leftover = 0
padding = EFFECTIVE_SEGMENT - leftover if leftover > 0 else 0
target = real_frame_len + padding
num_segments = (target - PREV_SEGMENT_COND_FRAMES + EFFECTIVE_SEGMENT - 1) // EFFECTIVE_SEGMENT
return max(1, num_segments), target
def zigzag_index(p, real_frame_len):
"""اندیس فریم واقعی برای موقعیت p (بعد از انتهای ویدیو، آینه‌ای برمی‌گردد)."""
if p < real_frame_len:
return p
period = 2 * real_frame_len
m = p % period
return m if m < real_frame_len else period - 1 - m
def pixels_to_uint8(pixels):
"""پیکسل‌های [-1, 1] (حاصل از فریم‌های ۸ بیتی) را بدون خطا به uint8 برمی‌گرداند (نصف حجم bf16)."""
return ((pixels.float() + 1.0) * 127.5).round().clamp(0, 255).to(torch.uint8)
def pixels_from_uint8(pixels_u8, device):
"""دقیقاً همان محاسبه‌ی normalize در diffusers: (x / 255) * 2 - 1 با float32."""
x = pixels_u8.to(device).to(torch.float32) / 255.0
return 2.0 * x - 1.0
def resolve_frame_size(image_height, image_width, area, mod_value=16):
aspect_ratio = image_height / image_width
height = int(math.sqrt(area * aspect_ratio)) // mod_value * mod_value
width = int(math.sqrt(area / aspect_ratio)) // mod_value * mod_value
crop_width = width if width / height < image_width / image_height else image_width * height // image_height
crop_height = height if width / height >= image_width / image_height else image_height * width // image_width
crop_top = (height - crop_height) // 2
crop_left = (width - crop_width) // 2
return height, width, (crop_top, crop_left, crop_height, crop_width)
def build_prev_cond_latents(vae_encode_fn, tail_frames, latent_h, latent_w, height, width, device):
"""معادل WanAnimate2SegmentPrevFramesStep (بدون نیمه‌ی تصویر مرجع)."""
num_frames = SEGMENT_FRAME_LENGTH + 1
if tail_frames is not None:
mask_len = PREV_SEGMENT_COND_FRAMES
prev = tail_frames.to(device, torch.float32) # [3, mask_len, H, W]
prev = F.interpolate(prev.permute(1, 0, 2, 3), size=(height, width), mode="bicubic").permute(1, 0, 2, 3)
cond_pixels = torch.cat(
[prev, torch.zeros(3, num_frames - mask_len - 1, height, width, device=device)], dim=1
)
else:
mask_len = 0
cond_pixels = torch.zeros(3, num_frames - 1, height, width, device=device)
lat = vae_encode_fn(cond_pixels.unsqueeze(0)).squeeze(0)
msk = get_i2v_mask(lat.shape[1], latent_h, latent_w, mask_len, device=device).to(lat.dtype)
return torch.cat([msk, lat], dim=0)