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