Team Ai
Modelpublic

diffusers/matrix-game-2-modular

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes14downloads
causal_model.py950 linesDownload Raw Back to transformer
1from .attention import attention2from .model import (3    MatrixGameWanRMSNorm,4    rope_apply,5    MatrixGameWanLayerNorm,6    MatrixGameWan_CROSSATTENTION_CLASSES,7    rope_params,8    MLPProj,9    sinusoidal_embedding_1d,10)11from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin12from torch.nn.attention.flex_attention import create_block_mask, flex_attention13from diffusers.configuration_utils import ConfigMixin, register_to_config14from torch.nn.attention.flex_attention import BlockMask15from diffusers.models.modeling_utils import ModelMixin16import torch.nn as nn17import torch18import math19import torch.distributed as dist20from .action_module import ActionModule21 22 23def causal_rope_apply(x, grid_sizes, freqs, start_frame=0):24    n, c = x.size(2), x.size(3) // 225 26    # split freqs27    freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)28 29    # loop over samples30    output = []31    f, h, w = grid_sizes.tolist()32 33    for i in range(len(x)):34        seq_len = f * h * w35 36        # precompute multipliers37        x_i = torch.view_as_complex(38            x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)39        )40        freqs_i = torch.cat(41            [42                freqs[0][start_frame : start_frame + f]43                .view(f, 1, 1, -1)44                .expand(f, h, w, -1),45                freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),46                freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),47            ],48            dim=-1,49        ).reshape(seq_len, 1, -1)50 51        # apply rotary embedding52        x_i = torch.view_as_real(x_i * freqs_i).flatten(2)53        x_i = torch.cat([x_i, x[i, seq_len:]])54 55        # append to collection56        output.append(x_i)57    return torch.stack(output).type_as(x)58 59 60class MatrixGameWanCausalSelfAttention(nn.Module):61    def __init__(62        self, dim, num_heads, local_attn_size=-1, sink_size=0, qk_norm=True, eps=1e-663    ):64        assert dim % num_heads == 065        super().__init__()66        self.dim = dim67        self.num_heads = num_heads68        self.head_dim = dim // num_heads69        self.local_attn_size = local_attn_size70        self.sink_size = sink_size71        self.qk_norm = qk_norm72        self.eps = eps73        self.max_attention_size = (74            15 * 1 * 880 if local_attn_size == -1 else local_attn_size * 88075        )76        # layers77        self.q = nn.Linear(dim, dim)78        self.k = nn.Linear(dim, dim)79        self.v = nn.Linear(dim, dim)80        self.o = nn.Linear(dim, dim)81        self.norm_q = MatrixGameWanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()82        self.norm_k = MatrixGameWanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()83 84    def forward(85        self,86        x,87        seq_lens,88        grid_sizes,89        freqs,90        block_mask,91        kv_cache=None,92        current_start=0,93        cache_start=None,94    ):95        r"""96        Args:97            x(Tensor): Shape [B, L, C] # num_heads, C / num_heads]98            seq_lens(Tensor): Shape [B]99            grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)100            freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]101            block_mask (BlockMask)102        """103        b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim104        if cache_start is None:105            cache_start = current_start106 107        # query, key, value function108        def qkv_fn(x):109            q = self.norm_q(self.q(x)).view(b, s, n, d)110            k = self.norm_k(self.k(x)).view(b, s, n, d)111            v = self.v(x).view(b, s, n, d)112            return q, k, v113 114        q, k, v = qkv_fn(x)  # B, F, HW, C115 116        if kv_cache is None:117            roped_query = rope_apply(q, grid_sizes, freqs).type_as(v)118            roped_key = rope_apply(k, grid_sizes, freqs).type_as(v)119 120            padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]121            padded_roped_query = torch.cat(122                [123                    roped_query,124                    torch.zeros(125                        [q.shape[0], padded_length, q.shape[2], q.shape[3]],126                        device=q.device,127                        dtype=v.dtype,128                    ),129                ],130                dim=1,131            )132 133            padded_roped_key = torch.cat(134                [135                    roped_key,136                    torch.zeros(137                        [k.shape[0], padded_length, k.shape[2], k.shape[3]],138                        device=k.device,139                        dtype=v.dtype,140                    ),141                ],142                dim=1,143            )144 145            padded_v = torch.cat(146                [147                    v,148                    torch.zeros(149                        [v.shape[0], padded_length, v.shape[2], v.shape[3]],150                        device=v.device,151                        dtype=v.dtype,152                    ),153                ],154                dim=1,155            )156 157            x = flex_attention(158                query=padded_roped_query.transpose(2, 1),  # after: B, HW, F, C159                key=padded_roped_key.transpose(2, 1),160                value=padded_v.transpose(2, 1),161                block_mask=block_mask,162            )[:, :, :-padded_length].transpose(2, 1)163        else:164            assert grid_sizes.ndim == 1165            frame_seqlen = math.prod(grid_sizes[1:]).item()166            current_start_frame = current_start // frame_seqlen167            roped_query = causal_rope_apply(168                q, grid_sizes, freqs, start_frame=current_start_frame169            ).type_as(v)170            roped_key = causal_rope_apply(171                k, grid_sizes, freqs, start_frame=current_start_frame172            ).type_as(v)173 174            current_end = current_start + roped_query.shape[1]175            sink_tokens = self.sink_size * frame_seqlen176 177            kv_cache_size = kv_cache["k"].shape[1]178            num_new_tokens = roped_query.shape[1]179 180            if (current_end > kv_cache["global_end_index"].item()) and (181                num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size182            ):183                num_evicted_tokens = (184                    num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size185                )186                num_rolled_tokens = (187                    kv_cache["local_end_index"].item()188                    - num_evicted_tokens189                    - sink_tokens190                )191                kv_cache["k"][:, sink_tokens : sink_tokens + num_rolled_tokens] = (192                    kv_cache["k"][193                        :,194                        sink_tokens + num_evicted_tokens : sink_tokens195                        + num_evicted_tokens196                        + num_rolled_tokens,197                    ].clone()198                )199                kv_cache["v"][:, sink_tokens : sink_tokens + num_rolled_tokens] = (200                    kv_cache["v"][201                        :,202                        sink_tokens + num_evicted_tokens : sink_tokens203                        + num_evicted_tokens204                        + num_rolled_tokens,205                    ].clone()206                )207                # Insert the new keys/values at the end208                local_end_index = (209                    kv_cache["local_end_index"].item()210                    + current_end211                    - kv_cache["global_end_index"].item()212                    - num_evicted_tokens213                )214                local_start_index = local_end_index - num_new_tokens215                kv_cache["k"][:, local_start_index:local_end_index] = roped_key216                kv_cache["v"][:, local_start_index:local_end_index] = v217            else:218                # Assign new keys/values directly up to current_end219                local_end_index = (220                    kv_cache["local_end_index"].item()221                    + current_end222                    - kv_cache["global_end_index"].item()223                )224                local_start_index = local_end_index - num_new_tokens225 226                kv_cache["k"][:, local_start_index:local_end_index] = roped_key227                kv_cache["v"][:, local_start_index:local_end_index] = v228            x = attention(229                roped_query,230                kv_cache["k"][231                    :,232                    max(0, local_end_index - self.max_attention_size) : local_end_index,233                ],234                kv_cache["v"][235                    :,236                    max(0, local_end_index - self.max_attention_size) : local_end_index,237                ],238            )239            kv_cache["global_end_index"].fill_(current_end)240            kv_cache["local_end_index"].fill_(local_end_index)241 242        # output243        x = x.flatten(2)244        x = self.o(x)245        return x246 247 248class MatrixGameWanCausalAttentionBlock(nn.Module):249    def __init__(250        self,251        cross_attn_type,252        dim,253        ffn_dim,254        num_heads,255        local_attn_size=-1,256        sink_size=0,257        qk_norm=True,258        cross_attn_norm=False,259        action_config={},260        block_idx=0,261        eps=1e-6,262    ):263        super().__init__()264        self.dim = dim265        self.ffn_dim = ffn_dim266        self.num_heads = num_heads267        self.local_attn_size = local_attn_size268        self.qk_norm = qk_norm269        self.cross_attn_norm = cross_attn_norm270        self.eps = eps271        if len(action_config) != 0 and block_idx in action_config["blocks"]:272            self.action_model = ActionModule(273                **action_config, local_attn_size=self.local_attn_size274            )275        else:276            self.action_model = None277        # layers278        self.norm1 = MatrixGameWanLayerNorm(dim, eps)279        self.self_attn = MatrixGameWanCausalSelfAttention(280            dim, num_heads, local_attn_size, sink_size, qk_norm, eps281        )282        self.norm3 = (283            MatrixGameWanLayerNorm(dim, eps, elementwise_affine=True)284            if cross_attn_norm285            else nn.Identity()286        )287        self.cross_attn = MatrixGameWan_CROSSATTENTION_CLASSES[cross_attn_type](288            dim, num_heads, (-1, -1), qk_norm, eps289        )290        self.norm2 = MatrixGameWanLayerNorm(dim, eps)291        self.ffn = nn.Sequential(292            nn.Linear(dim, ffn_dim),293            nn.GELU(approximate="tanh"),294            nn.Linear(ffn_dim, dim),295        )296 297        # modulation298        self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)299 300    def forward(301        self,302        x,303        e,304        seq_lens,305        grid_sizes,306        freqs,307        context,308        block_mask,309        block_mask_mouse,310        block_mask_keyboard,311        num_frame_per_block=3,312        use_rope_keyboard=False,313        mouse_cond=None,314        keyboard_cond=None,315        kv_cache=None,316        kv_cache_mouse=None,317        kv_cache_keyboard=None,318        crossattn_cache=None,319        current_start=0,320        cache_start=None,321        context_lens=None,322    ):323        r"""324        Args:325            x(Tensor): Shape [B, L, C]326            e(Tensor): Shape [B, F, 6, C]327            seq_lens(Tensor): Shape [B], length of each sequence in batch328            grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)329            freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]330        """331        assert e.ndim == 4332        num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]333 334        e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)335 336        y = self.self_attn(337            (338                self.norm1(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))339                * (1 + e[1])340                + e[0]341            ).flatten(1, 2),342            seq_lens,343            grid_sizes,344            freqs,345            block_mask,346            kv_cache,347            current_start,348            cache_start,349        )350 351        x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(352            1, 2353        )354 355        # cross-attention & ffn function356        def cross_attn_ffn(357            x,358            context,359            e,360            mouse_cond,361            keyboard_cond,362            block_mask_mouse,363            block_mask_keyboard,364            kv_cache_mouse=None,365            kv_cache_keyboard=None,366            crossattn_cache=None,367            start_frame=0,368            use_rope_keyboard=False,369            num_frame_per_block=3,370        ):371            x = x + self.cross_attn(372                self.norm3(x.to(context.dtype)),373                context,374                crossattn_cache=crossattn_cache,375            )376            if self.action_model is not None:377                assert mouse_cond is not None or keyboard_cond is not None378                x = self.action_model(379                    x.to(context.dtype),380                    grid_sizes[0],381                    grid_sizes[1],382                    grid_sizes[2],383                    mouse_cond,384                    keyboard_cond,385                    block_mask_mouse,386                    block_mask_keyboard,387                    is_causal=True,388                    kv_cache_mouse=kv_cache_mouse,389                    kv_cache_keyboard=kv_cache_keyboard,390                    start_frame=start_frame,391                    use_rope_keyboard=use_rope_keyboard,392                    num_frame_per_block=num_frame_per_block,393                )394 395            y = self.ffn(396                (397                    self.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))398                    * (1 + e[4])399                    + e[3]400                ).flatten(1, 2)401            )402 403            x = x + (404                y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]405            ).flatten(1, 2)406            return x407 408        assert grid_sizes.ndim == 1409        x = cross_attn_ffn(410            x,411            context,412            e,413            mouse_cond,414            keyboard_cond,415            block_mask_mouse,416            block_mask_keyboard,417            kv_cache_mouse,418            kv_cache_keyboard,419            crossattn_cache,420            start_frame=current_start // math.prod(grid_sizes[1:]).item(),421            use_rope_keyboard=use_rope_keyboard,422            num_frame_per_block=num_frame_per_block,423        )424        return x425 426 427class CausalHead(nn.Module):428    def __init__(self, dim, out_dim, patch_size, eps=1e-6):429        super().__init__()430        self.dim = dim431        self.out_dim = out_dim432        self.patch_size = patch_size433        self.eps = eps434 435        # layers436        out_dim = math.prod(patch_size) * out_dim437        self.norm = MatrixGameWanLayerNorm(dim, eps)438        self.head = nn.Linear(dim, out_dim)439 440        # modulation441        self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)442 443    def forward(self, x, e):444        r"""445        Args:446            x(Tensor): Shape [B, L1, C]447            e(Tensor): Shape [B, F, 1, C]448        """449 450        num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]451        e = (self.modulation.unsqueeze(1) + e).chunk(2, dim=2)452        x = self.head(453            self.norm(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1])454            + e[0]455        )456        return x457 458 459class MatrixGameWanCausalModel(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin):460    r"""461    MatrixGameWan diffusion backbone supporting both text-to-video and image-to-video.462    """463 464    ignore_for_config = ["patch_size", "cross_attn_norm", "qk_norm", "text_dim"]465    _no_split_modules = ["MatrixGameWanAttentionBlock"]466    _supports_gradient_checkpointing = True467 468    @register_to_config469    def __init__(470        self,471        model_type="t2v",472        patch_size=(1, 2, 2),473        text_len=512,474        in_dim=36,475        dim=1536,476        ffn_dim=8960,477        freq_dim=256,478        text_dim=4096,479        out_dim=16,480        num_heads=12,481        num_layers=30,482        local_attn_size=-1,483        sink_size=0,484        qk_norm=True,485        cross_attn_norm=True,486        action_config={},487        eps=1e-6,488    ):489        r"""490        Initialize the diffusion model backbone.491 492        Args:493            model_type (`str`, *optional*, defaults to 't2v'):494                Model variant - 't2v' (text-to-video) or 'i2v' (image-to-video)495            patch_size (`tuple`, *optional*, defaults to (1, 2, 2)):496                3D patch dimensions for video embedding (t_patch, h_patch, w_patch)497            text_len (`int`, *optional*, defaults to 512):498                Fixed length for text embeddings499            in_dim (`int`, *optional*, defaults to 16):500                Input video channels (C_in)501            dim (`int`, *optional*, defaults to 2048):502                Hidden dimension of the transformer503            ffn_dim (`int`, *optional*, defaults to 8192):504                Intermediate dimension in feed-forward network505            freq_dim (`int`, *optional*, defaults to 256):506                Dimension for sinusoidal time embeddings507            text_dim (`int`, *optional*, defaults to 4096):508                Input dimension for text embeddings509            out_dim (`int`, *optional*, defaults to 16):510                Output video channels (C_out)511            num_heads (`int`, *optional*, defaults to 16):512                Number of attention heads513            num_layers (`int`, *optional*, defaults to 32):514                Number of transformer blocks515            local_attn_size (`int`, *optional*, defaults to -1):516                Window size for temporal local attention (-1 indicates global attention)517            sink_size (`int`, *optional*, defaults to 0):518                Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache519            qk_norm (`bool`, *optional*, defaults to True):520                Enable query/key normalization521            cross_attn_norm (`bool`, *optional*, defaults to False):522                Enable cross-attention normalization523            eps (`float`, *optional*, defaults to 1e-6):524                Epsilon value for normalization layers525        """526 527        super().__init__()528 529        assert model_type in ["i2v"]530        self.model_type = model_type531        self.use_action_module = len(action_config) > 0532        self.patch_size = patch_size533        self.text_len = text_len534        self.in_dim = in_dim535        self.dim = dim536        self.ffn_dim = ffn_dim537        self.freq_dim = freq_dim538        self.text_dim = text_dim539        self.out_dim = out_dim540        self.num_heads = num_heads541        self.num_layers = num_layers542        self.local_attn_size = local_attn_size543        self.qk_norm = qk_norm544        self.cross_attn_norm = cross_attn_norm545        self.eps = eps546 547        # embeddings548        self.patch_embedding = nn.Conv3d(549            in_dim, dim, kernel_size=patch_size, stride=patch_size550        )551 552        self.time_embedding = nn.Sequential(553            nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim)554        )555        self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))556 557        # blocks558        cross_attn_type = "i2v_cross_attn"559        self.blocks = nn.ModuleList(560            [561                MatrixGameWanCausalAttentionBlock(562                    cross_attn_type,563                    dim,564                    ffn_dim,565                    num_heads,566                    local_attn_size,567                    sink_size,568                    qk_norm,569                    cross_attn_norm,570                    action_config=action_config,571                    eps=eps,572                    block_idx=idx,573                )574                for idx in range(num_layers)575            ]576        )577 578        # head579        self.head = CausalHead(dim, out_dim, patch_size, eps)580 581        # buffers (don't use register_buffer otherwise dtype will be changed in to())582        assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0583        d = dim // num_heads584        self.freqs = torch.cat(585            [586                rope_params(1024, d - 4 * (d // 6)),587                rope_params(1024, 2 * (d // 6)),588                rope_params(1024, 2 * (d // 6)),589            ],590            dim=1,591        )592 593        if model_type == "i2v":594            self.img_emb = MLPProj(1280, dim)595 596        self.gradient_checkpointing = False597 598        self.block_mask = None599        self.block_mask_keyboard = None600        self.block_mask_mouse = None601        self.use_rope_keyboard = True602 603    def _set_gradient_checkpointing(self, module, value=False):604        self.gradient_checkpointing = value605 606    @staticmethod607    def _prepare_blockwise_causal_attn_mask(608        device: torch.device | str,609        num_frames: int = 9,610        frame_seqlen: int = 880,611        num_frame_per_block=1,612        local_attn_size=-1,613    ) -> BlockMask:614        """615        we will divide the token sequence into the following format616        [1 latent frame] [1 latent frame] ... [1 latent frame]617        We use flexattention to construct the attention mask618        """619        total_length = num_frames * frame_seqlen620 621        # we do right padding to get to a multiple of 128622        padded_length = math.ceil(total_length / 128) * 128 - total_length623 624        ends = torch.zeros(625            total_length + padded_length, device=device, dtype=torch.long626        )627 628        # Block-wise causal mask will attend to all elements that are before the end of the current chunk629        frame_indices = torch.arange(630            start=0,631            end=total_length,632            step=frame_seqlen * num_frame_per_block,633            device=device,634        )635 636        for tmp in frame_indices:637            ends[tmp : tmp + frame_seqlen * num_frame_per_block] = (638                tmp + frame_seqlen * num_frame_per_block639            )640 641        def attention_mask(b, h, q_idx, kv_idx):642            if local_attn_size == -1:643                return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)644            else:645                return (646                    (kv_idx < ends[q_idx])647                    & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))648                ) | (q_idx == kv_idx)649            # return ((kv_idx < total_length) & (q_idx < total_length))  | (q_idx == kv_idx) # bidirectional mask650 651        block_mask = create_block_mask(652            attention_mask,653            B=None,654            H=None,655            Q_LEN=total_length + padded_length,656            KV_LEN=total_length + padded_length,657            _compile=False,658            device=device,659        )660 661        import torch.distributed as dist662 663        if not dist.is_initialized() or dist.get_rank() == 0:664            print(665                f" cache a block wise causal mask with block size of {num_frame_per_block} frames"666            )667 668        return block_mask669 670    @staticmethod671    def _prepare_blockwise_causal_attn_mask_keyboard(672        device: torch.device | str,673        num_frames: int = 9,674        frame_seqlen: int = 880,675        num_frame_per_block=1,676        local_attn_size=-1,677    ) -> BlockMask:678        """679        we will divide the token sequence into the following format680        [1 latent frame] [1 latent frame] ... [1 latent frame]681        We use flexattention to construct the attention mask682        """683        total_length2 = num_frames * frame_seqlen684 685        # we do right padding to get to a multiple of 128686        padded_length2 = math.ceil(total_length2 / 32) * 32 - total_length2687        padded_length_kv2 = math.ceil(num_frames / 32) * 32 - num_frames688        ends2 = torch.zeros(689            total_length2 + padded_length2, device=device, dtype=torch.long690        )691 692        # Block-wise causal mask will attend to all elements that are before the end of the current chunk693        frame_indices2 = torch.arange(694            start=0,695            end=total_length2,696            step=frame_seqlen * num_frame_per_block,697            device=device,698        )699        cnt = num_frame_per_block700        for tmp in frame_indices2:701            ends2[tmp : tmp + frame_seqlen * num_frame_per_block] = cnt702            cnt += num_frame_per_block703 704        def attention_mask2(b, h, q_idx, kv_idx):705            if local_attn_size == -1:706                return (kv_idx < ends2[q_idx]) | (q_idx == kv_idx)707            else:708                return (709                    (kv_idx < ends2[q_idx])710                    & (kv_idx >= (ends2[q_idx] - local_attn_size))711                ) | (q_idx == kv_idx)712            # return ((kv_idx < total_length) & (q_idx < total_length))  | (q_idx == kv_idx) # bidirectional mask713 714        block_mask2 = create_block_mask(715            attention_mask2,716            B=None,717            H=None,718            Q_LEN=total_length2 + padded_length2,719            KV_LEN=num_frames + padded_length_kv2,720            _compile=False,721            device=device,722        )723 724        import torch.distributed as dist725 726        if not dist.is_initialized() or dist.get_rank() == 0:727            print(728                f" cache a block wise causal mask with block size of {num_frame_per_block} frames"729            )730 731        return block_mask2732 733    @staticmethod734    def _prepare_blockwise_causal_attn_mask_action(735        device: torch.device | str,736        num_frames: int = 9,737        frame_seqlen: int = 1,738        num_frame_per_block=1,739        local_attn_size=-1,740    ) -> BlockMask:741        """742        we will divide the token sequence into the following format743        [1 latent frame] [1 latent frame] ... [1 latent frame]744        We use flexattention to construct the attention mask745        """746        total_length2 = num_frames * frame_seqlen747 748        # we do right padding to get to a multiple of 128749        padded_length2 = math.ceil(total_length2 / 32) * 32 - total_length2750        padded_length_kv2 = math.ceil(num_frames / 32) * 32 - num_frames751        ends2 = torch.zeros(752            total_length2 + padded_length2, device=device, dtype=torch.long753        )754 755        # Block-wise causal mask will attend to all elements that are before the end of the current chunk756        frame_indices2 = torch.arange(757            start=0,758            end=total_length2,759            step=frame_seqlen * num_frame_per_block,760            device=device,761        )762        cnt = num_frame_per_block763        for tmp in frame_indices2:764            ends2[tmp : tmp + frame_seqlen * num_frame_per_block] = cnt765            cnt += num_frame_per_block766 767        def attention_mask2(b, h, q_idx, kv_idx):768            if local_attn_size == -1:769                return (kv_idx < ends2[q_idx]) | (q_idx == kv_idx)770            else:771                return (772                    (kv_idx < ends2[q_idx])773                    & (kv_idx >= (ends2[q_idx] - local_attn_size))774                ) | (q_idx == kv_idx)775            # return ((kv_idx < total_length) & (q_idx < total_length))  | (q_idx == kv_idx) # bidirectional mask776 777        block_mask2 = create_block_mask(778            attention_mask2,779            B=None,780            H=None,781            Q_LEN=total_length2 + padded_length2,782            KV_LEN=num_frames + padded_length_kv2,783            _compile=False,784            device=device,785        )786 787        import torch.distributed as dist788 789        if not dist.is_initialized() or dist.get_rank() == 0:790            print(791                f" cache a block wise causal mask with block size of {num_frame_per_block} frames"792            )793 794        return block_mask2795 796    def _forward_inference(797        self,798        x,799        t,800        visual_context,801        cond_concat,802        mouse_cond=None,803        keyboard_cond=None,804        kv_cache: dict = None,805        kv_cache_mouse=None,806        kv_cache_keyboard=None,807        crossattn_cache: dict = None,808        current_start: int = 0,809        cache_start: int = 0,810        num_frames_per_block=3,811    ):812        r"""813        Run the diffusion model with kv caching.814        See Algorithm 2 of CausVid paper https://arxiv.org/abs/2412.07772 for details.815        This function will be run for num_frame times.816        Process the latent frames one by one (1560 tokens each)817 818        Args:819            x (List[Tensor]):820                List of input video tensors, each with shape [C_in, F, H, W]821            t (Tensor):822                Diffusion timesteps tensor of shape [B]823            context (List[Tensor]):824                List of text embeddings each with shape [L, C]825            seq_len (`int`):826                Maximum sequence length for positional encoding827            clip_fea (Tensor, *optional*):828                CLIP image features for image-to-video mode829            y (List[Tensor], *optional*):830                Conditional video inputs for image-to-video mode, same shape as x831 832        Returns:833            List[Tensor]:834                List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]835        """836 837        if mouse_cond is not None or keyboard_cond is not None:838            assert self.use_action_module == True839        # params840        device = self.patch_embedding.weight.device841        if self.freqs.device != device:842            self.freqs = self.freqs.to(device)843 844        x = torch.cat([x, cond_concat], dim=1)  # B C' F H W845 846        # embeddings847        x = self.patch_embedding(x)848        grid_sizes = torch.tensor(x.shape[2:], dtype=torch.long)849 850        x = x.flatten(2).transpose(1, 2)  # B FHW C'851        seq_lens = torch.tensor([u.size(0) for u in x], dtype=torch.long)852        assert seq_lens[0] <= 15 * 1 * 880853 854        e = self.time_embedding(855            sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x)856        )857        e0 = (858            self.time_projection(e)859            .unflatten(1, (6, self.dim))860            .unflatten(dim=0, sizes=t.shape)861        )862        # context863        context_lens = None864        context = self.img_emb(visual_context)865        # arguments866        kwargs = dict(867            e=e0,868            seq_lens=seq_lens,869            grid_sizes=grid_sizes,870            freqs=self.freqs,871            context=context,872            mouse_cond=mouse_cond,873            context_lens=context_lens,874            keyboard_cond=keyboard_cond,875            block_mask=self.block_mask,876            block_mask_mouse=self.block_mask_mouse,877            block_mask_keyboard=self.block_mask_keyboard,878            use_rope_keyboard=self.use_rope_keyboard,879            num_frame_per_block=num_frames_per_block,880        )881 882        def create_custom_forward(module):883            def custom_forward(*inputs, **kwargs):884                return module(*inputs, **kwargs)885 886            return custom_forward887 888        for block_index, block in enumerate(self.blocks):889            if torch.is_grad_enabled() and self.gradient_checkpointing:890                kwargs.update(891                    {892                        "kv_cache": kv_cache[block_index],893                        "kv_cache_mouse": kv_cache_mouse[block_index],894                        "kv_cache_keyboard": kv_cache_keyboard[block_index],895                        "current_start": current_start,896                        "cache_start": cache_start,897                    }898                )899                x = torch.utils.checkpoint.checkpoint(900                    create_custom_forward(block),901                    x,902                    **kwargs,903                    use_reentrant=False,904                )905            else:906                kwargs.update(907                    {908                        "kv_cache": kv_cache[block_index],909                        "kv_cache_mouse": kv_cache_mouse[block_index],910                        "kv_cache_keyboard": kv_cache_keyboard[block_index],911                        "crossattn_cache": crossattn_cache[block_index],912                        "current_start": current_start,913                        "cache_start": cache_start,914                    }915                )916                x = block(x, **kwargs)917 918        # head919        x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2))920        # unpatchify921        x = self.unpatchify(x, grid_sizes)922        return x923 924    def forward(self, *args, **kwargs):925        return self._forward_inference(*args, **kwargs)926 927    def unpatchify(self, x, grid_sizes):928        r"""929        Reconstruct video tensors from patch embeddings.930 931        Args:932            x (List[Tensor]):933                List of patchified features, each with shape [L, C_out * prod(patch_size)]934            grid_sizes (Tensor):935                Original spatial-temporal grid dimensions before patching,936                    shape [3] (3 dimensions correspond to F_patches, H_patches, W_patches)937 938        Returns:939            List[Tensor]:940                Reconstructed video tensors with shape [C_out, F, H / 8, W / 8]941        """942 943        c = self.out_dim944        bs = x.shape[0]945        x = x.view(bs, *grid_sizes, *self.patch_size, c)946        x = torch.einsum("bfhwpqrc->bcfphqwr", x)947        x = x.reshape(bs, c, *[i * j for i, j in zip(grid_sizes, self.patch_size)])948        return x949 950