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)