Spaces:
Runtime error
Runtime error
File size: 20,922 Bytes
a8d7e70 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 | # ==========================================================================
# موتور استنتاج 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)
|