Team Ai
Modelpublic

diffusers/matrix-game-2-modular

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes14downloads
model.py782 linesDownload Raw Back to transformer
1# Copyright 2024-2025 The Alibaba MatrixGameWan Team Authors. All rights reserved.2import math3import numpy as np4import torch5import torch.amp as amp6import torch.nn as nn7from diffusers.configuration_utils import ConfigMixin, register_to_config8from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin9from diffusers.models.modeling_utils import ModelMixin10from einops import repeat, rearrange11from .action_module import ActionModule12from .attention import flash_attention13 14DISABLE_COMPILE = False  # get os env15__all__ = ["MatrixGameWanModel"]16 17 18def sinusoidal_embedding_1d(dim, position):19    # preprocess20    assert dim % 2 == 021    half = dim // 222    position = position.type(torch.float64)23 24    # calculation25    sinusoid = torch.outer(26        position, torch.pow(10000, -torch.arange(half).to(position).div(half))27    )28    x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)29    return x30 31 32# @amp.autocast(enabled=False)33def rope_params(max_seq_len, dim, theta=10000):34    assert dim % 2 == 035    freqs = torch.outer(36        torch.arange(max_seq_len),37        1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float64).div(dim)),38    )39    freqs = torch.polar(torch.ones_like(freqs), freqs)40    return freqs41 42 43# @amp.autocast(enabled=False)44def rope_apply(x, grid_sizes, freqs):45    n, c = x.size(2), x.size(3) // 246 47    # split freqs48    freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)49 50    # loop over samples51    output = []52    # print(grid_sizes.shape, len(grid_sizes.tolist()), grid_sizes.tolist()[0])53    f, h, w = grid_sizes.tolist()54    for i in range(len(x)):55        seq_len = f * h * w56 57        # precompute multipliers58        x_i = torch.view_as_complex(59            x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)60        )61        freqs_i = torch.cat(62            [63                freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),64                freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),65                freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),66            ],67            dim=-1,68        ).reshape(seq_len, 1, -1)69 70        # apply rotary embedding71        x_i = torch.view_as_real(x_i * freqs_i).flatten(2)72        x_i = torch.cat([x_i, x[i, seq_len:]])73 74        # append to collection75        output.append(x_i)76    return torch.stack(output).type_as(x)77 78 79class MatrixGameWanRMSNorm(nn.Module):80    def __init__(self, dim, eps=1e-5):81        super().__init__()82        self.dim = dim83        self.eps = eps84        self.weight = nn.Parameter(torch.ones(dim))85 86    def forward(self, x):87        r"""88        Args:89            x(Tensor): Shape [B, L, C]90        """91        return self._norm(x.float()).type_as(x) * self.weight92 93    def _norm(self, x):94        return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)95 96 97class MatrixGameWanLayerNorm(nn.LayerNorm):98    def __init__(self, dim, eps=1e-6, elementwise_affine=False):99        super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)100 101    def forward(self, x):102        r"""103        Args:104            x(Tensor): Shape [B, L, C]105        """106        return super().forward(x).type_as(x)107 108 109class MatrixGameWanSelfAttention(nn.Module):110    def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6):111        assert dim % num_heads == 0112        super().__init__()113        self.dim = dim114        self.num_heads = num_heads115        self.head_dim = dim // num_heads116        self.window_size = window_size117        self.qk_norm = qk_norm118        self.eps = eps119 120        # layers121        self.q = nn.Linear(dim, dim)122        self.k = nn.Linear(dim, dim)123        self.v = nn.Linear(dim, dim)124        self.o = nn.Linear(dim, dim)125        self.norm_q = MatrixGameWanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()126        self.norm_k = MatrixGameWanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()127 128    def forward(self, x, seq_lens, grid_sizes, freqs):129        r"""130        Args:131            x(Tensor): Shape [B, L, num_heads, C / num_heads]132            seq_lens(Tensor): Shape [B]133            grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)134            freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]135        """136        b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim137 138        # query, key, value function139        def qkv_fn(x):140            q = self.norm_q(self.q(x)).view(b, s, n, d)141            k = self.norm_k(self.k(x)).view(b, s, n, d)142            v = self.v(x).view(b, s, n, d)143            return q, k, v144 145        q, k, v = qkv_fn(x)146        # print(k.shape, seq_lens)147        x = flash_attention(148            q=rope_apply(q, grid_sizes, freqs),149            k=rope_apply(k, grid_sizes, freqs),150            v=v,151            k_lens=seq_lens,152            window_size=self.window_size,153        )154 155        # output156        x = x.flatten(2)157        x = self.o(x)158        return x159 160 161# class MatrixGameWanT2VCrossAttention(MatrixGameWanSelfAttention):162 163#     def forward(self, x, context, context_lens, crossattn_cache=None):164#         r"""165#         Args:166#             x(Tensor): Shape [B, L1, C]167#             context(Tensor): Shape [B, L2, C]168#             context_lens(Tensor): Shape [B]169#             crossattn_cache (List[dict], *optional*): Contains the cached key and value tensors for context embedding.170#         """171#         b, n, d = x.size(0), self.num_heads, self.head_dim172 173#         # compute query, key, value174#         q = self.norm_q(self.q(x)).view(b, -1, n, d)175 176#         if crossattn_cache is not None:177#             if not crossattn_cache["is_init"]:178#                 crossattn_cache["is_init"] = True179#                 k = self.norm_k(self.k(context)).view(b, -1, n, d)180#                 v = self.v(context).view(b, -1, n, d)181#                 crossattn_cache["k"] = k182#                 crossattn_cache["v"] = v183#             else:184#                 k = crossattn_cache["k"]185#                 v = crossattn_cache["v"]186#         else:187#             k = self.norm_k(self.k(context)).view(b, -1, n, d)188#             v = self.v(context).view(b, -1, n, d)189 190#         # compute attention191#         x = flash_attention(q, k, v, k_lens=context_lens)192 193#         # output194#         x = x.flatten(2)195#         x = self.o(x)196#         return x197 198 199# class MatrixGameWanGanCrossAttention(MatrixGameWanSelfAttention):200 201#     def forward(self, x, context, crossattn_cache=None):202#         r"""203#         Args:204#             x(Tensor): Shape [B, L1, C]205#             context(Tensor): Shape [B, L2, C]206#             context_lens(Tensor): Shape [B]207#             crossattn_cache (List[dict], *optional*): Contains the cached key and value tensors for context embedding.208#         """209#         b, n, d = x.size(0), self.num_heads, self.head_dim210 211#         # compute query, key, value212#         qq = self.norm_q(self.q(context)).view(b, 1, -1, d)213 214#         kk = self.norm_k(self.k(x)).view(b, -1, n, d)215#         vv = self.v(x).view(b, -1, n, d)216 217#         # compute attention218#         x = flash_attention(qq, kk, vv)219 220#         # output221#         x = x.flatten(2)222#         x = self.o(x)223#         return x224 225 226class MatrixGameWanI2VCrossAttention(MatrixGameWanSelfAttention):227    def forward(self, x, context, crossattn_cache=None):228        r"""229        Args:230            x(Tensor): Shape [B, L1, C]231            context(Tensor): Shape [B, L2, C]232            context_lens(Tensor): Shape [B]233        """234        b, n, d = x.size(0), self.num_heads, self.head_dim235 236        # compute query, key, value237        q = self.norm_q(self.q(x)).view(b, -1, n, d)238        if crossattn_cache is not None:239            if not crossattn_cache["is_init"]:240                crossattn_cache["is_init"] = True241                k = self.norm_k(self.k(context)).view(b, -1, n, d)242                v = self.v(context).view(b, -1, n, d)243                crossattn_cache["k"] = k244                crossattn_cache["v"] = v245            else:246                k = crossattn_cache["k"]247                v = crossattn_cache["v"]248        else:249            k = self.norm_k(self.k(context)).view(b, -1, n, d)250            v = self.v(context).view(b, -1, n, d)251        # compute attention252        x = flash_attention(q, k, v, k_lens=None)253 254        # output255        x = x.flatten(2)256        x = self.o(x)257        return x258 259 260MatrixGameWan_CROSSATTENTION_CLASSES = {261    "i2v_cross_attn": MatrixGameWanI2VCrossAttention,262}263 264 265def mul_add(x, y, z):266    return x.float() + y.float() * z.float()267 268 269def mul_add_add(x, y, z):270    return x.float() * (1 + y) + z271 272 273class MatrixGameWanAttentionBlock(nn.Module):274    def __init__(275        self,276        cross_attn_type,277        dim,278        ffn_dim,279        num_heads,280        window_size=(-1, -1),281        qk_norm=True,282        cross_attn_norm=False,283        action_config={},284        eps=1e-6,285    ):286        super().__init__()287        self.dim = dim288        self.ffn_dim = ffn_dim289        self.num_heads = num_heads290        self.window_size = window_size291        self.qk_norm = qk_norm292        self.cross_attn_norm = cross_attn_norm293        self.eps = eps294        if len(action_config) != 0:295            self.action_model = ActionModule(**action_config)296        else:297            self.action_model = None298        # layers299        self.norm1 = MatrixGameWanLayerNorm(dim, eps)300        self.self_attn = MatrixGameWanSelfAttention(dim, num_heads, window_size, qk_norm, eps)301        self.norm3 = (302            MatrixGameWanLayerNorm(dim, eps, elementwise_affine=True)303            if cross_attn_norm304            else nn.Identity()305        )306        self.cross_attn = MatrixGameWan_CROSSATTENTION_CLASSES[cross_attn_type](307            dim, num_heads, (-1, -1), qk_norm, eps308        )309        self.norm2 = MatrixGameWanLayerNorm(dim, eps)310        self.ffn = nn.Sequential(311            nn.Linear(dim, ffn_dim),312            nn.GELU(approximate="tanh"),313            nn.Linear(ffn_dim, dim),314        )315 316        # modulation317        self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)318 319    def forward(320        self,321        x,322        e,323        seq_lens,324        grid_sizes,325        freqs,326        context,327        mouse_cond=None,328        keyboard_cond=None,329        # context_lens,330    ):331        r"""332        Args:333            x(Tensor): Shape [B, L, C]334            e(Tensor): Shape [B, 6, C]335            seq_lens(Tensor): Shape [B], length of each sequence in batch336            grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)337            freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]338        """339        # assert e.dtype == torch.float32340        if e.dim() == 3:341            modulation = self.modulation342            # with amp.autocast(dtype=torch.float32):343            e = (self.modulation + e).chunk(6, dim=1)344        elif e.dim() == 4:345            modulation = self.modulation.unsqueeze(2)  # 1, 6, 1, dim346            # with amp.autocast("cuda", dtype=torch.float32):347            e = (modulation + e).chunk(6, dim=1)348            e = [ei.squeeze(1) for ei in e]349        # assert e[0].dtype == torch.float32350 351        # self-attention352        y = self.self_attn(353            self.norm1(x) * (1 + e[1]) + e[0], seq_lens, grid_sizes, freqs354        )355        # with amp.autocast(dtype=torch.float32):356        x = x + y * e[2]357 358        # cross-attention & ffn function359        def cross_attn_ffn(x, context, e, mouse_cond, keyboard_cond):360            dtype = context.dtype361            x = x + self.cross_attn(self.norm3(x.to(dtype)), context)362            if self.action_model is not None:363                assert mouse_cond is not None or keyboard_cond is not None364                x = self.action_model(365                    x.to(dtype),366                    grid_sizes[0],367                    grid_sizes[1],368                    grid_sizes[2],369                    mouse_cond,370                    keyboard_cond,371                )372            y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3])373            # with amp.autocast(dtype=torch.float32):374            x = x + y * e[5]375            return x376 377        x = cross_attn_ffn(x, context, e, mouse_cond, keyboard_cond)378        return x379 380 381class Head(nn.Module):382    def __init__(self, dim, out_dim, patch_size, eps=1e-6):383        super().__init__()384        self.dim = dim385        self.out_dim = out_dim386        self.patch_size = patch_size387        self.eps = eps388 389        # layers390        out_dim = math.prod(patch_size) * out_dim391        self.norm = MatrixGameWanLayerNorm(dim, eps)392        self.head = nn.Linear(dim, out_dim)393 394        # modulation395        self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)396 397    def forward(self, x, e):398        r"""399        Args:400            x(Tensor): Shape [B, L1, C]401            e(Tensor): Shape [B, C]402        """403        # assert e.dtype == torch.float32404        # with amp.autocast(dtype=torch.float32):405        if e.dim() == 2:406            modulation = self.modulation  # 1, 2, dim407            e = (modulation + e.unsqueeze(1)).chunk(2, dim=1)408        elif e.dim() == 3:409            modulation = self.modulation.unsqueeze(2)  # 1, 2, seq, dim410            e = (modulation + e.unsqueeze(1)).chunk(2, dim=1)411            e = [ei.squeeze(1) for ei in e]412        x = self.head(self.norm(x) * (1 + e[1]) + e[0])413        return x414 415 416class MLPProj(torch.nn.Module):417    def __init__(self, in_dim, out_dim):418        super().__init__()419 420        self.proj = torch.nn.Sequential(421            torch.nn.LayerNorm(in_dim),422            torch.nn.Linear(in_dim, in_dim),423            torch.nn.GELU(),424            torch.nn.Linear(in_dim, out_dim),425            torch.nn.LayerNorm(out_dim),426        )427 428    def forward(self, image_embeds):429        clip_extra_context_tokens = self.proj(image_embeds)430        return clip_extra_context_tokens431 432 433# class RegisterTokens(nn.Module):434#     def __init__(self, num_registers: int, dim: int):435#         super().__init__()436#         self.register_tokens = nn.Parameter(torch.randn(num_registers, dim) * 0.02)437#         self.rms_norm = MatrixGameWanRMSNorm(dim, eps=1e-6)438 439#     def forward(self):440#         return self.rms_norm(self.register_tokens)441 442#     def reset_parameters(self):443#         nn.init.normal_(self.register_tokens, std=0.02)444 445 446class MatrixGameWanModel(ModelMixin, ConfigMixin, FromOriginalModelMixin, PeftAdapterMixin):447    r"""448    MatrixGameWan diffusion backbone supporting both text-to-video and image-to-video.449    """450 451    ignore_for_config = [452        "patch_size",453        "cross_attn_norm",454        "qk_norm",455        "text_dim",456        "window_size",457    ]458    _no_split_modules = ["MatrixGameWanAttentionBlock"]459    _supports_gradient_checkpointing = True460 461    @register_to_config462    def __init__(463        self,464        model_type="i2v",465        patch_size=(1, 2, 2),466        text_len=512,467        in_dim=36,468        dim=1536,469        ffn_dim=8960,470        freq_dim=256,471        text_dim=4096,472        out_dim=16,473        num_heads=12,474        num_layers=30,475        window_size=(-1, -1),476        qk_norm=True,477        cross_attn_norm=True,478        inject_sample_info=False,479        action_config={},480        eps=1e-6,481    ):482        r"""483        Initialize the diffusion model backbone.484 485        Args:486            model_type (`str`, *optional*, defaults to 't2v'):487                Model variant - 't2v' (text-to-video) or 'i2v' (image-to-video)488            patch_size (`tuple`, *optional*, defaults to (1, 2, 2)):489                3D patch dimensions for video embedding (t_patch, h_patch, w_patch)490            text_len (`int`, *optional*, defaults to 512):491                Fixed length for text embeddings492            in_dim (`int`, *optional*, defaults to 16):493                Input video channels (C_in)494            dim (`int`, *optional*, defaults to 2048):495                Hidden dimension of the transformer496            ffn_dim (`int`, *optional*, defaults to 8192):497                Intermediate dimension in feed-forward network498            freq_dim (`int`, *optional*, defaults to 256):499                Dimension for sinusoidal time embeddings500            text_dim (`int`, *optional*, defaults to 4096):501                Input dimension for text embeddings502            out_dim (`int`, *optional*, defaults to 16):503                Output video channels (C_out)504            num_heads (`int`, *optional*, defaults to 16):505                Number of attention heads506            num_layers (`int`, *optional*, defaults to 32):507                Number of transformer blocks508            window_size (`tuple`, *optional*, defaults to (-1, -1)):509                Window size for local attention (-1 indicates global attention)510            qk_norm (`bool`, *optional*, defaults to True):511                Enable query/key normalization512            cross_attn_norm (`bool`, *optional*, defaults to False):513                Enable cross-attention normalization514            eps (`float`, *optional*, defaults to 1e-6):515                Epsilon value for normalization layers516        """517 518        super().__init__()519 520        assert model_type in ["i2v"]521        self.model_type = model_type522        self.use_action_module = len(action_config) > 0523        assert self.use_action_module == True524        self.patch_size = patch_size525        self.text_len = text_len526        self.in_dim = in_dim527        self.dim = dim528        self.ffn_dim = ffn_dim529        self.freq_dim = freq_dim530        self.text_dim = text_dim531        self.out_dim = out_dim532        self.num_heads = num_heads533        self.num_layers = num_layers534        self.window_size = window_size535        self.qk_norm = qk_norm536        self.cross_attn_norm = cross_attn_norm537        self.eps = eps538        self.local_attn_size = -1539 540        # embeddings541        self.patch_embedding = nn.Conv3d(542            in_dim, dim, kernel_size=patch_size, stride=patch_size543        )544        # self.text_embedding = nn.Sequential(545        #     nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),546        #     nn.Linear(dim, dim))547 548        self.time_embedding = nn.Sequential(549            nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim)550        )551        self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))552 553        # blocks554        cross_attn_type = "i2v_cross_attn"555        self.blocks = nn.ModuleList(556            [557                MatrixGameWanAttentionBlock(558                    cross_attn_type,559                    dim,560                    ffn_dim,561                    num_heads,562                    window_size,563                    qk_norm,564                    cross_attn_norm,565                    eps=eps,566                    action_config=action_config,567                )568                for _ in range(num_layers)569            ]570        )571 572        # head573        self.head = Head(dim, out_dim, patch_size, eps)574 575        # buffers (don't use register_buffer otherwise dtype will be changed in to())576        assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0577        d = dim // num_heads578        self.freqs = torch.cat(579            [580                rope_params(1024, d - 4 * (d // 6)),581                rope_params(1024, 2 * (d // 6)),582                rope_params(1024, 2 * (d // 6)),583            ],584            dim=1,585        )586 587        if model_type == "i2v":588            self.img_emb = MLPProj(1280, dim)589 590        # initialize weights591        self.init_weights()592 593        self.gradient_checkpointing = False594 595    def _set_gradient_checkpointing(self, module, value=False):596        self.gradient_checkpointing = value597 598    def forward(self, *args, **kwargs):599        # if kwargs.get('classify_mode', False) is True:600        # kwargs.pop('classify_mode')601        # return self._forward_classify(*args, **kwargs)602        # else:603        return self._forward(*args, **kwargs)604 605    def _forward(606        self,607        x,608        t,609        visual_context,610        cond_concat,611        mouse_cond=None,612        keyboard_cond=None,613        fps=None,614        # seq_len,615        # classify_mode=False,616        # concat_time_embeddings=False,617        # register_tokens=None,618        # cls_pred_branch=None,619        # gan_ca_blocks=None,620        # clip_fea=None,621        # y=None,622    ):623        r"""624        Forward pass through the diffusion model625 626        Args:627            x (List[Tensor]):628                List of input video tensors, each with shape [C_in, F, H, W]629            t (Tensor):630                Diffusion timesteps tensor of shape [B]631            context (List[Tensor]):632                List of text embeddings each with shape [L, C]633            seq_len (`int`):634                Maximum sequence length for positional encoding635            clip_fea (Tensor, *optional*):636                CLIP image features for image-to-video mode637            y (List[Tensor], *optional*):638                Conditional video inputs for image-to-video mode, same shape as x639 640        Returns:641            List[Tensor]:642                List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]643        """644        # params645        if mouse_cond is not None or keyboard_cond is not None:646            assert self.use_action_module == True647        device = self.patch_embedding.weight.device648        if self.freqs.device != device:649            self.freqs = self.freqs.to(device)650 651        x = torch.cat([x, cond_concat], dim=1)652        # embeddings653        x = self.patch_embedding(x)654        grid_sizes = torch.tensor(x.shape[2:], dtype=torch.long)655        x = x.flatten(2).transpose(1, 2)656        seq_lens = torch.tensor([u.size(0) for u in x], dtype=torch.long)657        # seq_len = seq_lens.max()658        # # assert seq_lens.max() <= seq_len659        # x = torch.cat([660        #     torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],661        #               dim=1) for u in x662        # ])663 664        # time embeddings665        # with amp.autocast(dtype=torch.float32):666        # assert t.ndim == 1667        e = self.time_embedding(668            sinusoidal_embedding_1d(self.freq_dim, t).type_as(x)669        )  # TODO: check if t ndim == 1670 671        e0 = self.time_projection(e).unflatten(1, (6, self.dim))672        # assert e.dtype == torch.float32 and e0.dtype == torch.float32673 674        # context675        context_lens = None676        # context = self.text_embedding(677        #     torch.stack([678        #         torch.cat(679        #             [u, u.new_zeros(self.text_len - u.size(0), u.size(1))])680        #         for u in context681        #     ]))682 683        # if clip_fea is not None:684        #     context_clip = self.img_emb(clip_fea)  # bs x 257 x dim685        context = self.img_emb(visual_context)686 687        # arguments688        # kwargs = dict(689        #     e=e0,690        #     seq_lens=seq_lens,691        #     grid_sizes=grid_sizes,692        #     freqs=self.freqs,693        #     context=context,694        #     context_lens=context_lens)695        kwargs = dict(696            e=e0,697            grid_sizes=grid_sizes,698            seq_lens=seq_lens,699            freqs=self.freqs,700            context=context,701            mouse_cond=mouse_cond,702            # context_lens=context_lens,703            keyboard_cond=keyboard_cond,704        )705 706        def create_custom_forward(module):707            def custom_forward(*inputs, **kwargs):708                return module(*inputs, **kwargs)709 710            return custom_forward711 712        for ii, block in enumerate(self.blocks):713            if torch.is_grad_enabled() and self.gradient_checkpointing:714                x = torch.utils.checkpoint.checkpoint(715                    create_custom_forward(block),716                    x,717                    **kwargs,718                    use_reentrant=False,719                )720            else:721                x = block(x, **kwargs)722 723        # head724        x = self.head(x, e)725 726        # unpatchify727        x = self.unpatchify(x, grid_sizes)728 729        return x.float()730 731    def unpatchify(self, x, grid_sizes):  # TODO check grid sizes732        r"""733        Reconstruct video tensors from patch embeddings.734 735        Args:736            x (List[Tensor]):737                List of patchified features, each with shape [L, C_out * prod(patch_size)]738            grid_sizes (Tensor):739                Original spatial-temporal grid dimensions before patching,740                    shape [3] (3 dimensions correspond to F_patches, H_patches, W_patches)741 742        Returns:743            List[Tensor]:744                Reconstructed video tensors with shape [C_out, F, H / 8, W / 8]745        """746 747        c = self.out_dim748        bs = x.shape[0]749        x = x.view(bs, *grid_sizes, *self.patch_size, c)750        x = torch.einsum("bfhwpqrc->bcfphqwr", x)751        x = x.reshape(bs, c, *[i * j for i, j in zip(grid_sizes, self.patch_size)])752        return x753 754    def init_weights(self):755        r"""756        Initialize model parameters using Xavier initialization.757        """758 759        # basic init760        for m in self.modules():761            if isinstance(m, nn.Linear):762                nn.init.xavier_uniform_(m.weight)763                if m.bias is not None:764                    nn.init.zeros_(m.bias)765 766        # init embeddings767        nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))768        for m in self.time_embedding.modules():769            if isinstance(m, nn.Linear):770                nn.init.normal_(m.weight, std=0.02)771 772        # init output layer773        nn.init.zeros_(self.head.head.weight)774        if self.use_action_module == True:775            for m in self.blocks:776                nn.init.zeros_(m.action_model.proj_mouse.weight)777                if m.action_model.proj_mouse.bias is not None:778                    nn.init.zeros_(m.action_model.proj_mouse.bias)779                nn.init.zeros_(m.action_model.proj_keyboard.weight)780                if m.action_model.proj_keyboard.bias is not None:781                    nn.init.zeros_(m.action_model.proj_keyboard.bias)782