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