diffusers/matrix-game-2-modular
014
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 