Spaces:
Runtime error
Runtime error
Download animate2_engine.py from Opera8/Wan2.2-Animate2: direct link, hf CLI and curl.
- Browser
- Download file 20.9 kB
-
https://huggingface.co/spaces/Opera8/Wan2.2-Animate2/resolve/main/animate2_engine.py
- Command line
-
hf download hf://spaces/Opera8/Wan2.2-Animate2/animate2_engine.py
-
curl -L -o animate2_engine.py https://huggingface.co/spaces/Opera8/Wan2.2-Animate2/resolve/main/animate2_engine.py
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 | |
| 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 بدون پدینگ (خروجی توکنهای واقعی با نسخهی پدشده یکسان است) | |
| # -------------------------------------------------------------------------- | |
| 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) | |