# ========================================================================== # موتور استنتاج 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)