Team Ai
Modelpublic

diffusers/matrix-game-2-modular

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes14downloads
action_module.py1149 linesDownload Raw Back to transformer
1from typing import Any, List, Tuple, Optional, Union, Dict2from einops import rearrange3from flash_attn import flash_attn_func4import torch5import torch.nn as nn6import math7from torch.nn.attention.flex_attention import flex_attention8 9try:10    import flash_attn11 12except:13    from flash_attn import flash_attn_func14 15FLASH_ATTN_3_AVAILABLE = False16 17 18DISABLE_COMPILE = False  # get os env19flex_attention = torch.compile(20    flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs"21)22 23import torch24from typing import Union, Tuple, List25 26 27def _to_tuple(x, dim=2):28    if isinstance(x, int):29        return (x,) * dim30    elif len(x) == dim:31        return x32    else:33        raise ValueError(f"Expected length {dim} or int, but got {x}")34 35 36def get_meshgrid_nd(start, *args, dim=2):37    """38    Get n-D meshgrid with start, stop and num.39 40    Args:41        start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,42            step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num43            should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in44            n-tuples.45        *args: See above.46        dim (int): Dimension of the meshgrid. Defaults to 2.47 48    Returns:49        grid (np.ndarray): [dim, ...]50    """51    if len(args) == 0:52        # start is grid_size53        num = _to_tuple(start, dim=dim)54        start = (0,) * dim55        stop = num56    elif len(args) == 1:57        # start is start, args[0] is stop, step is 158        start = _to_tuple(start, dim=dim)59        stop = _to_tuple(args[0], dim=dim)60        num = [stop[i] - start[i] for i in range(dim)]61    elif len(args) == 2:62        # start is start, args[0] is stop, args[1] is num63        start = _to_tuple(start, dim=dim)  # Left-Top       eg: 12,064        stop = _to_tuple(args[0], dim=dim)  # Right-Bottom   eg: 20,3265        num = _to_tuple(args[1], dim=dim)  # Target Size    eg: 32,12466    else:67        raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")68 69    # PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)70    axis_grid = []71    for i in range(dim):72        a, b, n = start[i], stop[i], num[i]73        g = torch.linspace(a, b, n + 1, dtype=torch.float32, device=torch.cuda.current_device())[:n]74        axis_grid.append(g)75    grid = torch.meshgrid(*axis_grid, indexing="ij")  # dim x [W, H, D]76    grid = torch.stack(grid, dim=0)  # [dim, W, H, D]77 78    return grid79 80 81#################################################################################82#                   Rotary Positional Embedding Functions                       #83#################################################################################84# https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L8085 86 87def reshape_for_broadcast(88    freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],89    x: torch.Tensor,90    head_first=False,91):92    """93    Reshape frequency tensor for broadcasting it with another tensor.94 95    This function reshapes the frequency tensor to have the same shape as the target tensor 'x'96    for the purpose of broadcasting the frequency tensor during element-wise operations.97 98    Notes:99        When using FlashMHAModified, head_first should be False.100        When using Attention, head_first should be True.101 102    Args:103        freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.104        x (torch.Tensor): Target tensor for broadcasting compatibility.105        head_first (bool): head dimension first (except batch dim) or not.106 107    Returns:108        torch.Tensor: Reshaped frequency tensor.109 110    Raises:111        AssertionError: If the frequency tensor doesn't match the expected shape.112        AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.113    """114    ndim = x.ndim115    assert 0 <= 1 < ndim116 117    if isinstance(freqs_cis, tuple):118        # freqs_cis: (cos, sin) in real space119        if head_first:120            assert freqs_cis[0].shape == (121                x.shape[-2],122                x.shape[-1],123            ), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"124            shape = [125                d if i == ndim - 2 or i == ndim - 1 else 1126                for i, d in enumerate(x.shape)127            ]128        else:129            # assert freqs_cis[0].shape == (130            #     x.shape[1],131            #     x.shape[-1],132            # ), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"133            # shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]134            shape = [1, freqs_cis[0].shape[0], 1, freqs_cis[0].shape[1]]135        return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)136    else:137        # freqs_cis: values in complex space138        if head_first:139            assert freqs_cis.shape == (140                x.shape[-2],141                x.shape[-1],142            ), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"143            shape = [144                d if i == ndim - 2 or i == ndim - 1 else 1145                for i, d in enumerate(x.shape)146            ]147        else:148            assert freqs_cis.shape == (149                x.shape[1],150                x.shape[-1],151            ), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"152            shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]153        return freqs_cis.view(*shape)154 155 156def rotate_half(x):157    x_real, x_imag = (158        x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)159    )  # [B, S, H, D//2]160    return torch.stack([-x_imag, x_real], dim=-1).flatten(3)161 162 163def apply_rotary_emb(164    xq: torch.Tensor,165    xk: torch.Tensor,166    freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],167    head_first: bool = False,168    start_offset: int = 0,169) -> Tuple[torch.Tensor, torch.Tensor]:170    """171    Apply rotary embeddings to input tensors using the given frequency tensor.172 173    This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided174    frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor175    is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are176    returned as real tensors.177 178    Args:179        xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]180        xk (torch.Tensor): Key tensor to apply rotary embeddings.   [B, S, H, D]181        freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential.182        head_first (bool): head dimension first (except batch dim) or not.183 184    Returns:185        Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.186 187    """188    # print(freqs_cis[0].shape, xq.shape, xk.shape)189    xk_out = None190    assert isinstance(freqs_cis, tuple)191    if isinstance(freqs_cis, tuple):192        cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first)  # [S, D]193        cos, sin = cos.to(xq.device), sin.to(xq.device)194        # real * cos - imag * sin195        # imag * cos + real * sin196        xq_out = (xq.float() * cos[:, start_offset:start_offset + xq.shape[1], :, :] + rotate_half(xq.float()) * sin[:, start_offset:start_offset + xq.shape[1], :, :]).type_as(xq)197        xk_out = (xk.float() * cos[:, start_offset:start_offset + xk.shape[1], :, :] + rotate_half(xk.float()) * sin[:, start_offset:start_offset + xk.shape[1], :, :]).type_as(xk)198    else:199        # view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)200        xq_ = torch.view_as_complex(201            xq.float().reshape(*xq.shape[:-1], -1, 2)202        )  # [B, S, H, D//2]203        freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(204            xq.device205        )  # [S, D//2] --> [1, S, 1, D//2]206        # (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)207        # view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)208        xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)209        xk_ = torch.view_as_complex(210            xk.float().reshape(*xk.shape[:-1], -1, 2)211        )  # [B, S, H, D//2]212        xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)213 214    return xq_out, xk_out215 216 217def get_nd_rotary_pos_embed(218    rope_dim_list,219    start,220    *args,221    theta=10000.0,222    use_real=False,223    theta_rescale_factor: Union[float, List[float]] = 1.0,224    interpolation_factor: Union[float, List[float]] = 1.0,225):226    """227    This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.228 229    Args:230        rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.231            sum(rope_dim_list) should equal to head_dim of attention layer.232        start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,233            args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.234        *args: See above.235        theta (float): Scaling factor for frequency computation. Defaults to 10000.0.236        use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.237            Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real238            part and an imaginary part separately.239        theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.240 241    Returns:242        pos_embed (torch.Tensor): [HW, D/2]243    """244 245    grid = get_meshgrid_nd(246        start, *args, dim=len(rope_dim_list)247    )  # [3, W, H, D] / [2, W, H]248 249    if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):250        theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)251    elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:252        theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)253    assert len(theta_rescale_factor) == len(254        rope_dim_list255    ), "len(theta_rescale_factor) should equal to len(rope_dim_list)"256 257    if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):258        interpolation_factor = [interpolation_factor] * len(rope_dim_list)259    elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:260        interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)261    assert len(interpolation_factor) == len(262        rope_dim_list263    ), "len(interpolation_factor) should equal to len(rope_dim_list)"264 265    # use 1/ndim of dimensions to encode grid_axis266    embs = []267    for i in range(len(rope_dim_list)):268        emb = get_1d_rotary_pos_embed(269            rope_dim_list[i],270            grid[i].reshape(-1),271            theta,272            use_real=use_real,273            theta_rescale_factor=theta_rescale_factor[i],274            interpolation_factor=interpolation_factor[i],275        )  # 2 x [WHD, rope_dim_list[i]]276        embs.append(emb)277 278    if use_real:279        cos = torch.cat([emb[0] for emb in embs], dim=1)  # (WHD, D/2)280        sin = torch.cat([emb[1] for emb in embs], dim=1)  # (WHD, D/2)281        return cos, sin282    else:283        emb = torch.cat(embs, dim=1)  # (WHD, D/2)284        return emb285 286 287def get_1d_rotary_pos_embed(288    dim: int,289    pos: Union[torch.FloatTensor, int],290    theta: float = 10000.0,291    use_real: bool = False,292    theta_rescale_factor: float = 1.0,293    interpolation_factor: float = 1.0,294) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:295    """296    Precompute the frequency tensor for complex exponential (cis) with given dimensions.297    (Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)298 299    This function calculates a frequency tensor with complex exponential using the given dimension 'dim'300    and the end index 'end'. The 'theta' parameter scales the frequencies.301    The returned tensor contains complex values in complex64 data type.302 303    Args:304        dim (int): Dimension of the frequency tensor.305        pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar306        theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.307        use_real (bool, optional): If True, return real part and imaginary part separately.308                                   Otherwise, return complex numbers.309        theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.310 311    Returns:312        freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]313        freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]314    """315    if isinstance(pos, int):316        pos = torch.arange(pos, device=torch.cuda.current_device()).float()317 318    # proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning319    # has some connection to NTK literature320    if theta_rescale_factor != 1.0:321        theta *= theta_rescale_factor ** (dim / (dim - 2))322 323    freqs = 1.0 / (324        theta ** (torch.arange(0, dim, 2, device=torch.cuda.current_device())[: (dim // 2)].float() / dim)325    )  # [D/2]326    # assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"327    freqs = torch.outer(pos * interpolation_factor, freqs)  # [S, D/2]328    if use_real:329        freqs_cos = freqs.cos().repeat_interleave(2, dim=1)  # [S, D]330        freqs_sin = freqs.sin().repeat_interleave(2, dim=1)  # [S, D]331        return freqs_cos, freqs_sin332    else:333        freqs_cis = torch.polar(334            torch.ones_like(freqs), freqs335        )  # complex64     # [S, D/2]336        return freqs_cis337 338 339class MatrixGameWanRMSNorm(nn.Module):340    def __init__(self, dim, eps=1e-5):341        super().__init__()342        self.dim = dim343        self.eps = eps344        self.weight = nn.Parameter(torch.ones(dim))345 346    def forward(self, x):347        r"""348        Args:349            x(Tensor): Shape [B, L, C]350        """351        return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)352 353 354class ActionModule(nn.Module):355    """356    action module from https://arxiv.org/pdf/2501.08325357    ้ผ ๆ ‡ๆŽงๅˆถไฟกๅท็š„่พ“ๅ…ฅๆ˜ฏไธ€ไธช L*D ็š„ๅ‘้‡358    ้”ฎ็›˜ๅŒๆ ท359    """360 361    def __init__(362        self,363        mouse_dim_in: int = 2,364        keyboard_dim_in: int = 6,365        hidden_size: int = 128,366        img_hidden_size: int = 1536,367        keyboard_hidden_dim: int = 1024,368        mouse_hidden_dim: int = 1024,369        vae_time_compression_ratio: int = 4,370        windows_size: int = 3,371        heads_num: int = 16,372        patch_size: list = [1, 2, 2],373        qk_norm: bool = True,374        qkv_bias: bool = False,375        rope_dim_list: list = [8, 28, 28],376        rope_theta=256,377        mouse_qk_dim_list=[8, 28, 28],378        enable_mouse=True,379        enable_keyboard=True,380        local_attn_size=6,381        blocks=[],382    ):383        device = None384 385        super().__init__()386        self.local_attn_size = local_attn_size387        self.enable_mouse = enable_mouse388        self.enable_keyboard = enable_keyboard389 390        self.rope_dim_list = rope_dim_list391        self.rope_theta = rope_theta392        if self.enable_keyboard:393            self.keyboard_embed = nn.Sequential(394                nn.Linear(keyboard_dim_in, hidden_size, bias=True),395                nn.SiLU(),396                nn.Linear(hidden_size, hidden_size, bias=True),397            )398 399        self.mouse_qk_dim_list = mouse_qk_dim_list400        self.heads_num = heads_num401        if self.enable_mouse:402            c = mouse_hidden_dim403            self.mouse_mlp = torch.nn.Sequential(404                torch.nn.Linear(405                    mouse_dim_in * vae_time_compression_ratio * windows_size406                    + img_hidden_size,407                    c,408                    bias=True,409                ),410                torch.nn.GELU(approximate="tanh"),411                torch.nn.Linear(c, c),412                torch.nn.LayerNorm(c),413            )414 415            head_dim = c // heads_num416            self.t_qkv = nn.Linear(c, c * 3, bias=qkv_bias)417            self.img_attn_q_norm = (418                MatrixGameWanRMSNorm(head_dim, eps=1e-6) if qk_norm else nn.Identity()419            )420            self.img_attn_k_norm = (421                MatrixGameWanRMSNorm(head_dim, eps=1e-6) if qk_norm else nn.Identity()422            )423            self.proj_mouse = nn.Linear(c, img_hidden_size, bias=qkv_bias)424 425        if self.enable_keyboard:426            head_dim_key = keyboard_hidden_dim // heads_num427            self.key_attn_q_norm = (428                MatrixGameWanRMSNorm(head_dim_key, eps=1e-6) if qk_norm else nn.Identity()429            )430            self.key_attn_k_norm = (431                MatrixGameWanRMSNorm(head_dim_key, eps=1e-6) if qk_norm else nn.Identity()432            )433 434            self.mouse_attn_q = nn.Linear(435                img_hidden_size, keyboard_hidden_dim, bias=qkv_bias436            )437            self.keyboard_attn_kv = nn.Linear(438                hidden_size * windows_size * vae_time_compression_ratio,439                keyboard_hidden_dim * 2,440                bias=qkv_bias,441            )442            self.proj_keyboard = nn.Linear(443                keyboard_hidden_dim, img_hidden_size, bias=qkv_bias444            )445 446        self.vae_time_compression_ratio = vae_time_compression_ratio447        self.windows_size = windows_size448        self.patch_size = patch_size449        self.freqs_cos, self.freqs_sin = self.get_rotary_pos_embed(450            7500,451            self.patch_size[1],452            self.patch_size[2],453            64,454            self.mouse_qk_dim_list,455            start_offset=0,456        )457 458    def patchify(self, x, patch_size):459        """460        x : (N C T H W)461        """462        pt, ph, pw = self.patch_size463        t, h, w = x.shape[2] // pt, x.shape[3] // ph, x.shape[4] // pw464        c = x.shape[1]465        x = x.reshape(shape=(x.shape[0], c, t, pt, h, ph, w, pw))466        x = torch.einsum("nctohpwq->nthwcopq", x)467        x = x.reshape(shape=(x.shape[0], t * h * w, c * pt * ph * pw))468        return x469 470    def unpatchify(self, x, t, h, w, patch_size):471        """472        x: (N, T, patch_size**2 * C)473        imgs: (N, H, W, C)474        """475        c = x.shape[2] // patch_size  # self.unpatchify_channels476        pt, ph, pw = self.patch_size477        assert t * h * w == x.shape[1]478 479        x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))480        x = torch.einsum("nthwcopq->nctohpwq", x)481        imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))482 483        return imgs484 485    def get_rotary_pos_embed(486        self, video_length, height, width, head_dim, rope_dim_list=None, start_offset=0487    ):488        target_ndim = 3489        ndim = 5 - 2490 491        latents_size = [video_length + start_offset, height, width]492 493        if isinstance(self.patch_size, int):494            assert all(s % self.patch_size == 0 for s in latents_size), (495                f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "496                f"but got {latents_size}."497            )498            rope_sizes = [s // self.patch_size for s in latents_size]499        elif isinstance(self.patch_size, list):500            assert all(501                s % self.patch_size[idx] == 0 for idx, s in enumerate(latents_size)502            ), (503                f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "504                f"but got {latents_size}."505            )506            rope_sizes = [507                s // self.patch_size[idx] for idx, s in enumerate(latents_size)508            ]509 510        if len(rope_sizes) != target_ndim:511            rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes  # time axis512 513        if rope_dim_list is None:514            rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]515        assert (516            sum(rope_dim_list) == head_dim517        ), "sum(rope_dim_list) should equal to head_dim of attention layer"518        freqs_cos, freqs_sin = get_nd_rotary_pos_embed(519            rope_dim_list,520            rope_sizes,521            theta=self.rope_theta,522            use_real=True,523            theta_rescale_factor=1,524        )525        return freqs_cos[526            -video_length * rope_sizes[1] * rope_sizes[2] // self.patch_size[0] :527        ], freqs_sin[528            -video_length * rope_sizes[1] * rope_sizes[2] // self.patch_size[0] :529        ]530 531    def forward(532        self,533        x,534        tt,535        th,536        tw,537        mouse_condition=None,538        keyboard_condition=None,539        block_mask_mouse=None,540        block_mask_keyboard=None,541        is_causal=False,542        kv_cache_mouse=None,543        kv_cache_keyboard=None,544        start_frame=0,545        use_rope_keyboard=True,546        num_frame_per_block=3,547    ):548        """549        hidden_states: B, tt*th*tw, C550        mouse_condition: B, N_frames, C1551        keyboard_condition: B, N_frames, C2552        """553        assert use_rope_keyboard == True554 555        B, N_frames, C = keyboard_condition.shape556 557        assert tt * th * tw == x.shape[1]558        assert (559            (N_frames - 1) + self.vae_time_compression_ratio560        ) % self.vae_time_compression_ratio == 0561        N_feats = int((N_frames - 1) / self.vae_time_compression_ratio) + 1562 563        # Defined freqs_cis early so it's available for both mouse and keyboard564        freqs_cis = (self.freqs_cos, self.freqs_sin)565 566        assert (567            N_feats == tt and ((is_causal and kv_cache_mouse == None) or not is_causal)568        ) or (569            (N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block and is_causal570        )571 572        if self.enable_mouse and mouse_condition is not None:573            hidden_states = rearrange(574                x, "B (T S) C -> (B S) T C", T=tt, S=th * tw575            )  # 65*272*480 -> 17*(272//16)*(480//16) -> 8670576            B, N_frames, C = mouse_condition.shape577        else:578            hidden_states = x579        # padding580 581        pad_t = self.vae_time_compression_ratio * self.windows_size582        if self.enable_mouse and mouse_condition is not None:583            pad = mouse_condition[:, 0:1, :].expand(-1, pad_t, -1)584            mouse_condition = torch.cat([pad, mouse_condition], dim=1)585            if is_causal and kv_cache_mouse is not None:586                mouse_condition = mouse_condition[587                    :,588                    self.vae_time_compression_ratio589                    * (N_feats - num_frame_per_block - self.windows_size)590                    + pad_t :,591                    :,592                ]593                group_mouse = [594                    mouse_condition[595                        :,596                        self.vae_time_compression_ratio * (i - self.windows_size)597                        + pad_t : i * self.vae_time_compression_ratio + pad_t,598                        :,599                    ]600                    for i in range(num_frame_per_block)601                ]602            else:603                group_mouse = [604                    mouse_condition[605                        :,606                        self.vae_time_compression_ratio * (i - self.windows_size)607                        + pad_t : i * self.vae_time_compression_ratio + pad_t,608                        :,609                    ]610                    for i in range(N_feats)611                ]612 613            group_mouse = torch.stack(group_mouse, dim=1)614 615            S = th * tw616            group_mouse = group_mouse.unsqueeze(-1).expand(617                B, num_frame_per_block, pad_t, C, S618            )619            group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(620                B * S, num_frame_per_block, pad_t * C621            )622 623            group_mouse = torch.cat([hidden_states, group_mouse], dim=-1)624            group_mouse = self.mouse_mlp(group_mouse)625 626            # qkv627            mouse_qkv = self.t_qkv(group_mouse)628            q, k, v = rearrange(629                mouse_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num630            )  # BHW F H C631            q = self.img_attn_q_norm(q).to(v)632            k = self.img_attn_k_norm(k).to(v)633            # rope embd634 635            # freqs_cis = (self.freqs_cos, self.freqs_sin)636 637            q, k = apply_rotary_emb(638                q, k, freqs_cis, start_offset=start_frame, head_first=False639            )640            ## TODO: adding cache here641            if is_causal:642                if kv_cache_mouse is None:643                    assert (644                        q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0645                    )  # == 880, f"{q.shape[0]},{k.shape[0]}"646                    padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]647                    padded_q = torch.cat(648                        [649                            q,650                            torch.zeros(651                                [q.shape[0], padded_length, q.shape[2], q.shape[3]],652                                device=q.device,653                                dtype=v.dtype,654                            ),655                        ],656                        dim=1,657                    )658                    padded_k = torch.cat(659                        [660                            k,661                            torch.zeros(662                                [k.shape[0], padded_length, k.shape[2], k.shape[3]],663                                device=k.device,664                                dtype=v.dtype,665                            ),666                        ],667                        dim=1,668                    )669                    padded_v = torch.cat(670                        [671                            v,672                            torch.zeros(673                                [v.shape[0], padded_length, v.shape[2], v.shape[3]],674                                device=v.device,675                                dtype=v.dtype,676                            ),677                        ],678                        dim=1,679                    )680                    attn = flex_attention(681                        query=padded_q.transpose(2, 1),  # after: B, HW, F, C682                        key=padded_k.transpose(2, 1),683                        value=padded_v.transpose(2, 1),684                        block_mask=block_mask_mouse,685                    )[:, :, :-padded_length].transpose(2, 1)686                else:687                    current_start = start_frame688                    current_end = current_start + q.shape[1]689 690                    assert q.shape[1] == num_frame_per_block691                    sink_size = 0692                    max_attention_size = self.local_attn_size693                    sink_tokens = sink_size * 1694                    kv_cache_size = kv_cache_mouse["k"].shape[1]695                    num_new_tokens = q.shape[1]696 697                    if (current_end > kv_cache_mouse["global_end_index"].item()) and (698                        num_new_tokens + kv_cache_mouse["local_end_index"].item()699                        > kv_cache_size700                    ):701                        num_evicted_tokens = (702                            num_new_tokens703                            + kv_cache_mouse["local_end_index"].item()704                            - kv_cache_size705                        )706                        num_rolled_tokens = (707                            kv_cache_mouse["local_end_index"].item()708                            - num_evicted_tokens709                            - sink_tokens710                        )711                        kv_cache_mouse["k"][712                            :, sink_tokens : sink_tokens + num_rolled_tokens713                        ] = kv_cache_mouse["k"][714                            :,715                            sink_tokens + num_evicted_tokens : sink_tokens716                            + num_evicted_tokens717                            + num_rolled_tokens,718                        ].clone()719                        kv_cache_mouse["v"][720                            :, sink_tokens : sink_tokens + num_rolled_tokens721                        ] = kv_cache_mouse["v"][722                            :,723                            sink_tokens + num_evicted_tokens : sink_tokens724                            + num_evicted_tokens725                            + num_rolled_tokens,726                        ].clone()727                        # Insert the new keys/values at the end728                        local_end_index = (729                            kv_cache_mouse["local_end_index"].item()730                            + current_end731                            - kv_cache_mouse["global_end_index"].item()732                            - num_evicted_tokens733                        )734                        local_start_index = local_end_index - num_new_tokens735                    else:736                        local_end_index = (737                            kv_cache_mouse["local_end_index"].item()738                            + current_end739                            - kv_cache_mouse["global_end_index"].item()740                        )741                        local_start_index = local_end_index - num_new_tokens742 743                    kv_cache_mouse["k"][:, local_start_index:local_end_index] = k744                    kv_cache_mouse["v"][:, local_start_index:local_end_index] = v745 746                    if FLASH_ATTN_3_AVAILABLE:747                        attn, attn_prob = flash_attn.flash_attn_func(748                            q,749                            kv_cache_mouse["k"][750                                :,751                                max(752                                    0, local_end_index - max_attention_size753                                ) : local_end_index,754                            ],755                            kv_cache_mouse["v"][756                                :,757                                max(758                                    0, local_end_index - max_attention_size759                                ) : local_end_index,760                            ],761                        )762                    else:763                        attn = flash_attn_func(764                            q,765                            kv_cache_mouse["k"][766                                :,767                                max(768                                    0, local_end_index - max_attention_size769                                ) : local_end_index,770                            ],771                            kv_cache_mouse["v"][772                                :,773                                max(774                                    0, local_end_index - max_attention_size775                                ) : local_end_index,776                            ],777                        )778                    kv_cache_mouse["global_end_index"].fill_(current_end)779                    kv_cache_mouse["local_end_index"].fill_(local_end_index)780            else:781                attn = flash_attn_func(782                    q,  # 880, f, 16, 64783                    k,  # 880, f, 16, 64784                    v,  # 880, f, 16, 64785                )786            # Compute cu_squlens and max_seqlen for flash attention787            # qk norm788            attn = rearrange(attn, "(b S) T h d -> b (T S) (h d)", b=B)789 790            hidden_states = rearrange(x, "(B S) T C -> B (T S) C", B=B)791            attn = self.proj_mouse(attn)792 793            hidden_states = hidden_states + attn794 795        if self.enable_keyboard and keyboard_condition is not None:796            pad = keyboard_condition[:, 0:1, :].expand(-1, pad_t, -1)797            keyboard_condition = torch.cat([pad, keyboard_condition], dim=1)798            if is_causal and kv_cache_keyboard is not None:799                keyboard_condition = keyboard_condition[800                    :,801                    self.vae_time_compression_ratio802                    * (N_feats - num_frame_per_block - self.windows_size)803                    + pad_t :,804                    :,805                ]  # keyboard_condition[:, self.vae_time_compression_ratio*(start_frame - self.windows_size) + pad_t:start_frame * self.vae_time_compression_ratio + pad_t,:]806                keyboard_condition = self.keyboard_embed(keyboard_condition)807                group_keyboard = [808                    keyboard_condition[809                        :,810                        self.vae_time_compression_ratio * (i - self.windows_size)811                        + pad_t : i * self.vae_time_compression_ratio + pad_t,812                        :,813                    ]814                    for i in range(num_frame_per_block)815                ]816            else:817                keyboard_condition = self.keyboard_embed(keyboard_condition)818                group_keyboard = [819                    keyboard_condition[820                        :,821                        self.vae_time_compression_ratio * (i - self.windows_size)822                        + pad_t : i * self.vae_time_compression_ratio + pad_t,823                        :,824                    ]825                    for i in range(N_feats)826                ]827            group_keyboard = torch.stack(group_keyboard, dim=1)  # B F RW C828            group_keyboard = group_keyboard.reshape(829                shape=(group_keyboard.shape[0], group_keyboard.shape[1], -1)830            )831            # apply cross attn832            mouse_q = self.mouse_attn_q(hidden_states)833            keyboard_kv = self.keyboard_attn_kv(group_keyboard)834 835            B, L, HD = mouse_q.shape836            D = HD // self.heads_num837            q = mouse_q.view(B, L, self.heads_num, D)838 839            B, L, KHD = keyboard_kv.shape840            k, v = keyboard_kv.view(B, L, 2, self.heads_num, D).permute(2, 0, 1, 3, 4)841 842            # Compute cu_squlens and max_seqlen for flash attention843            # qk norm844 845            q = self.key_attn_q_norm(q).to(v)846            k = self.key_attn_k_norm(k).to(v)847            S = th * tw848            assert S == 880849            # position embed850            if use_rope_keyboard:851                B, TS, H, D = q.shape852                T_ = TS // S853                q = q.view(B, T_, S, H, D).transpose(1, 2).reshape(B * S, T_, H, D)854                q, k = apply_rotary_emb(855                    q, k, freqs_cis, start_offset=start_frame, head_first=False856                )857 858                k1, k2, k3, k4 = k.shape859                k = k.expand(S, k2, k3, k4)860                v = v.expand(S, k2, k3, k4)861 862                if is_causal:863                    if kv_cache_keyboard is None:864                        assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0865 866                        padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]867                        padded_q = torch.cat(868                            [869                                q,870                                torch.zeros(871                                    [q.shape[0], padded_length, q.shape[2], q.shape[3]],872                                    device=q.device,873                                    dtype=v.dtype,874                                ),875                            ],876                            dim=1,877                        )878                        padded_k = torch.cat(879                            [880                                k,881                                torch.zeros(882                                    [k.shape[0], padded_length, k.shape[2], k.shape[3]],883                                    device=k.device,884                                    dtype=v.dtype,885                                ),886                            ],887                            dim=1,888                        )889                        padded_v = torch.cat(890                            [891                                v,892                                torch.zeros(893                                    [v.shape[0], padded_length, v.shape[2], v.shape[3]],894                                    device=v.device,895                                    dtype=v.dtype,896                                ),897                            ],898                            dim=1,899                        )900                        attn = flex_attention(901                            query=padded_q.transpose(2, 1),  # after: B, HW, F, C902                            key=padded_k.transpose(2, 1),903                            value=padded_v.transpose(2, 1),904                            block_mask=block_mask_keyboard,905                        )[:, :, :-padded_length].transpose(2, 1)906                    else:907                        current_start = start_frame908                        current_end = current_start + k.shape[1]909                        assert k.shape[1] == num_frame_per_block910                        sink_size = 0911                        max_attention_size = self.local_attn_size912                        sink_tokens = sink_size * 1913                        kv_cache_size = kv_cache_keyboard["k"].shape[1]914                        num_new_tokens = k.shape[1]915 916                        if (917                            current_end > kv_cache_keyboard["global_end_index"].item()918                        ) and (919                            num_new_tokens + kv_cache_keyboard["local_end_index"].item()920                            > kv_cache_size921                        ):922                            num_evicted_tokens = (923                                num_new_tokens924                                + kv_cache_keyboard["local_end_index"].item()925                                - kv_cache_size926                            )927                            num_rolled_tokens = (928                                kv_cache_keyboard["local_end_index"].item()929                                - num_evicted_tokens930                                - sink_tokens931                            )932                            kv_cache_keyboard["k"][933                                :, sink_tokens : sink_tokens + num_rolled_tokens934                            ] = kv_cache_keyboard["k"][935                                :,936                                sink_tokens + num_evicted_tokens : sink_tokens937                                + num_evicted_tokens938                                + num_rolled_tokens,939                            ].clone()940                            kv_cache_keyboard["v"][941                                :, sink_tokens : sink_tokens + num_rolled_tokens942                            ] = kv_cache_keyboard["v"][943                                :,944                                sink_tokens + num_evicted_tokens : sink_tokens945                                + num_evicted_tokens946                                + num_rolled_tokens,947                            ].clone()948                            # Insert the new keys/values at the end949                            local_end_index = (950                                kv_cache_keyboard["local_end_index"].item()951                                + current_end952                                - kv_cache_keyboard["global_end_index"].item()953                                - num_evicted_tokens954                            )955                            local_start_index = local_end_index - num_new_tokens956                        else:957                            local_end_index = (958                                kv_cache_keyboard["local_end_index"].item()959                                + current_end960                                - kv_cache_keyboard["global_end_index"].item()961                            )962                            local_start_index = local_end_index - num_new_tokens963                        assert (964                            k.shape[0] == 880965                        )  # BS == 1 or the cache should not be saved/ load method should be modified966                        kv_cache_keyboard["k"][:, local_start_index:local_end_index] = (967                            k[:1]968                        )969                        kv_cache_keyboard["v"][:, local_start_index:local_end_index] = (970                            v[:1]971                        )972 973                        if FLASH_ATTN_3_AVAILABLE:974                            attn, attn_prob = flash_attn.flash_attn_func(975                                q,976                                kv_cache_keyboard["k"][977                                    :,978                                    max(979                                        0, local_end_index - max_attention_size980                                    ) : local_end_index,981                                ].repeat(S, 1, 1, 1),982                                kv_cache_keyboard["v"][983                                    :,984                                    max(985                                        0, local_end_index - max_attention_size986                                    ) : local_end_index,987                                ].repeat(S, 1, 1, 1),988                            )989                        else:990                            attn = flash_attn_func(991                                q,992                                kv_cache_keyboard["k"][993                                    :,994                                    max(995                                        0, local_end_index - max_attention_size996                                    ) : local_end_index,997                                ].repeat(S, 1, 1, 1),998                                kv_cache_keyboard["v"][999                                    :,1000                                    max(1001                                        0, local_end_index - max_attention_size1002                                    ) : local_end_index,1003                                ].repeat(S, 1, 1, 1),1004                            )1005 1006                        kv_cache_keyboard["global_end_index"].fill_(current_end)1007                        kv_cache_keyboard["local_end_index"].fill_(local_end_index)1008                else:1009                    attn = flash_attn_func(1010                        q,  # 1, f*880, 16, 641011                        k,  # 1, f, 16, 641012                        v,  # 1, f, 16, 641013                        causal=False,1014                    )1015                attn = rearrange(attn, "(B S) T H D -> B (T S) (H D)", S=S)1016            else:1017                if is_causal:1018                    if kv_cache_keyboard is None:1019                        padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]1020                        padded_q = torch.cat(1021                            [1022                                q,1023                                torch.zeros(1024                                    [q.shape[0], padded_length, q.shape[2], q.shape[3]],1025                                    device=q.device,1026                                    dtype=v.dtype,1027                                ),1028                            ],1029                            dim=1,1030                        )1031                        padded_k = torch.cat(1032                            [1033                                k,1034                                torch.zeros(1035                                    [k.shape[0], padded_length, k.shape[2], k.shape[3]],1036                                    device=k.device,1037                                    dtype=v.dtype,1038                                ),1039                            ],1040                            dim=1,1041                        )1042                        padded_v = torch.cat(1043                            [1044                                v,1045                                torch.zeros(1046                                    [v.shape[0], padded_length, v.shape[2], v.shape[3]],1047                                    device=v.device,1048                                    dtype=v.dtype,1049                                ),1050                            ],1051                            dim=1,1052                        )1053                        attn = flex_attention(1054                            query=padded_q.transpose(2, 1),  # after: B, HW, F, C1055                            key=padded_k.transpose(2, 1),1056                            value=padded_v.transpose(2, 1),1057                            block_mask=block_mask_keyboard,1058                        )[:, :, :-padded_length].transpose(2, 1)1059                    else:1060                        current_start = start_frame1061                        current_end = current_start + k.shape[1]1062                        assert k.shape[1] == num_frame_per_block1063                        sink_size = 01064                        local_attn_size = self.local_attn_size1065                        max_attention_size = self.local_attn_size1066                        sink_tokens = sink_size * 11067                        kv_cache_size = kv_cache_keyboard["k"].shape[1]1068                        num_new_tokens = k.shape[1]1069 1070                        if (1071                            current_end > kv_cache_keyboard["global_end_index"].item()1072                        ) and (1073                            num_new_tokens + kv_cache_keyboard["local_end_index"].item()1074                            > kv_cache_size1075                        ):1076                            num_evicted_tokens = (1077                                num_new_tokens1078                                + kv_cache_keyboard["local_end_index"].item()1079                                - kv_cache_size1080                            )1081                            num_rolled_tokens = (1082                                kv_cache_keyboard["local_end_index"].item()1083                                - num_evicted_tokens1084                                - sink_tokens1085                            )1086                            kv_cache_keyboard["k"][1087                                :, sink_tokens : sink_tokens + num_rolled_tokens1088                            ] = kv_cache_keyboard["k"][1089                                :,1090                                sink_tokens + num_evicted_tokens : sink_tokens1091                                + num_evicted_tokens1092                                + num_rolled_tokens,1093                            ].clone()1094                            kv_cache_keyboard["v"][1095                                :, sink_tokens : sink_tokens + num_rolled_tokens1096                            ] = kv_cache_keyboard["v"][1097                                :,1098                                sink_tokens + num_evicted_tokens : sink_tokens1099                                + num_evicted_tokens1100                                + num_rolled_tokens,1101                            ].clone()1102                            # Insert the new keys/values at the end1103                            local_end_index = (1104                                kv_cache_keyboard["local_end_index"].item()1105                                + current_end1106                                - kv_cache_keyboard["global_end_index"].item()1107                                - num_evicted_tokens1108                            )1109                            local_start_index = local_end_index - num_new_tokens1110 1111                        else:1112                            local_end_index = (1113                                kv_cache_keyboard["local_end_index"].item()1114                                + current_end1115                                - kv_cache_keyboard["global_end_index"].item()1116                            )1117                            local_start_index = local_end_index - num_new_tokens1118                        kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k1119                        kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v1120                        attn = flash_attn_func(1121                            q,1122                            kv_cache_keyboard["k"][1123                                :,1124                                max(1125                                    0, local_end_index - max_attention_size1126                                ) : local_end_index,1127                            ],1128                            kv_cache_keyboard["v"][1129                                :,1130                                max(1131                                    0, local_end_index - max_attention_size1132                                ) : local_end_index,1133                            ],1134                            # causal=is_causal1135                        )1136                        kv_cache_keyboard["global_end_index"].fill_(current_end)1137                        kv_cache_keyboard["local_end_index"].fill_(local_end_index)1138                else:1139                    attn = flash_attn_func(1140                        q,  # 1, f*880, 16, 641141                        k,  # 1, f, 16, 641142                        v,  # 1, f, 16, 641143                        # causal=is_causal,1144                    )1145                attn = rearrange(attn, "B L H D -> B L (H D)")1146            attn = self.proj_keyboard(attn)1147            hidden_states = hidden_states + attn1148        return hidden_states1149