From a21e4b3210f2fd137793b5d499974364c98f86be Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 28 Aug 2025 20:54:36 +0300 Subject: [PATCH] Update model.py --- wanvideo/modules/model.py | 46 +++++---------------------------------- 1 file changed, 6 insertions(+), 40 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 22f39df..4c17e30 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1119,7 +1119,7 @@ class WanAttentionBlock(nn.Module): x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs) x = x + x_motion * mtv_strength - if self.rope_func == "comfy_chunked": + if self.rope_func == "comfy_chunked" and not self.zero_timestep: y = self.ffn_chunked(x, shift_mlp, scale_mlp) else: norm2_x = self.norm2(x) @@ -1677,50 +1677,16 @@ class WanModel(torch.nn.Module): attention_mode=attention_mode ) self.trainable_cond_mask = nn.Embedding(3, self.dim) - self.adain_mode = adain_mode - self.zero_timestep = zero_timestep - - # init motioner - enable_framepack = False - enable_motioner = False - add_last_motion = False - if enable_motioner and enable_framepack: - raise ValueError( - "enable_motioner and enable_framepack are mutually exclusive, please set one of them to False" - ) - self.enable_motioner = enable_motioner - self.add_last_motion = add_last_motion - # if enable_motioner: - # motioner_dim = 2048 - # self.motioner = MotionerTransformers( - # patch_size=(2, 4, 4), - # dim=motioner_dim, - # ffn_dim=motioner_dim, - # freq_dim=256, - # out_dim=16, - # num_heads=16, - # num_layers=13, - # window_size=(-1, -1), - # qk_norm=True, - # cross_attn_norm=False, - # eps=1e-6, - # motion_token_num=4, - # enable_tsm=False, - # motion_stride=4, - # expand_ratio=2, - # trainable_token_pos_emb=False, - # ) - # self.zip_motion_out = torch.nn.Sequential( - # WanLayerNorm(motioner_dim), - # zero_module(nn.Linear(motioner_dim, self.dim))) - - enable_framepack = True - if enable_framepack: + self.frame_packer = FramePackMotioner( inner_dim=self.dim, num_heads=self.num_heads, zip_frame_buckets=[1, 2, 16], drop_mode='padd') + self.adain_mode = adain_mode + self.zero_timestep = zero_timestep + + @staticmethod def _prepare_blockwise_causal_attn_mask(