From f9be754980f5b7c2eae496fa91a15fb1f4d037ff Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 00:17:05 +0300 Subject: [PATCH] cleanup --- wanvideo/modules/model.py | 147 +------------------------------------- 1 file changed, 1 insertion(+), 146 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4c17e30..3bde4c8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1677,7 +1677,7 @@ class WanModel(torch.nn.Module): attention_mode=attention_mode ) self.trainable_cond_mask = nn.Embedding(3, self.dim) - + self.frame_packer = FramePackMotioner( inner_dim=self.dim, num_heads=self.num_heads, @@ -1838,151 +1838,6 @@ class WanModel(torch.nn.Module): block.to(self.offload_device, non_blocking=self.use_non_blocking) return hints - - def process_motion(self, motion_latents, drop_motion_frames=False): - if drop_motion_frames or motion_latents[0].shape[1] == 0: - return [], [] - self.lat_motion_frames = motion_latents[0].shape[1] - mot = [self.patch_embedding(m.unsqueeze(0)) for m in motion_latents] - batch_size = len(mot) - - mot_remb = [] - flattern_mot = [] - for bs in range(batch_size): - height, width = mot[bs].shape[3], mot[bs].shape[4] - flat_mot = mot[bs].flatten(2).transpose(1, 2).contiguous() - motion_grid_sizes = [[ - torch.tensor([-self.lat_motion_frames, 0, - 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([0, height, width]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.lat_motion_frames, height, - width]).unsqueeze(0).repeat(1, 1) - ]] - motion_rope_emb = rope_precompute( - flat_mot.detach().view(1, flat_mot.shape[1], self.num_heads, - self.dim // self.num_heads), - motion_grid_sizes, - self.freqs, - start=None) - mot_remb.append(motion_rope_emb) - flattern_mot.append(flat_mot) - return flattern_mot, mot_remb - - def process_motion_frame_pack(self, - motion_latents, - drop_motion_frames=False, - add_last_motion=2): - flattern_mot, mot_remb = self.frame_packer(motion_latents, - add_last_motion) - if drop_motion_frames: - return [m[:, :0] for m in flattern_mot - ], [m[:, :0] for m in mot_remb] - else: - return flattern_mot, mot_remb - - def process_motion_transformer_motioner(self, - motion_latents, - drop_motion_frames=False, - add_last_motion=True): - batch_size, height, width = len( - motion_latents), motion_latents[0].shape[2] // self.patch_size[ - 1], motion_latents[0].shape[3] // self.patch_size[2] - - freqs = self.freqs - device = self.patch_embedding.weight.device - if freqs.device != device: - freqs = freqs.to(device) - if self.trainable_token_pos_emb: - token_freqs = self.token_freqs.to(torch.float64) - token_freqs = token_freqs / token_freqs.norm( - dim=-1, keepdim=True) - freqs = [freqs, torch.view_as_complex(token_freqs)] - - if not drop_motion_frames and add_last_motion: - last_motion_latent = [u[:, -1:] for u in motion_latents] - last_mot = [ - self.patch_embedding(m.unsqueeze(0)) for m in last_motion_latent - ] - last_mot = [m.flatten(2).transpose(1, 2) for m in last_mot] - last_mot = torch.cat(last_mot) - gride_sizes = [[ - torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([0, height, - width]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([1, height, - width]).unsqueeze(0).repeat(batch_size, 1) - ]] - else: - last_mot = torch.zeros([batch_size, 0, self.dim], - device=motion_latents[0].device, - dtype=motion_latents[0].dtype) - gride_sizes = [] - - zip_motion = self.motioner(motion_latents) - zip_motion = self.zip_motion_out(zip_motion) - if drop_motion_frames: - zip_motion = zip_motion * 0.0 - zip_motion_grid_sizes = [[ - torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([ - 0, self.motioner.motion_side_len, self.motioner.motion_side_len - ]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor( - [1 if not self.trainable_token_pos_emb else -1, height, - width]).unsqueeze(0).repeat(batch_size, 1), - ]] - - mot = torch.cat([last_mot, zip_motion], dim=1) - gride_sizes = gride_sizes + zip_motion_grid_sizes - - motion_rope_emb = rope_precompute( - mot.detach().view(batch_size, mot.shape[1], self.num_heads, - self.dim // self.num_heads), - gride_sizes, - freqs, - start=None) - return [m.unsqueeze(0) for m in mot - ], [r.unsqueeze(0) for r in motion_rope_emb] - - def inject_motion(self, - x, - seq_lens, - rope_embs, - mask_input, - motion_latents, - drop_motion_frames=False, - add_last_motion=True): - # inject the motion frames token to the hidden states - if self.enable_motioner: - mot, mot_remb = self.process_motion_transformer_motioner( - motion_latents, - drop_motion_frames=drop_motion_frames, - add_last_motion=add_last_motion) - elif self.enable_framepack: - mot, mot_remb = self.process_motion_frame_pack( - motion_latents, - drop_motion_frames=drop_motion_frames, - add_last_motion=add_last_motion) - else: - mot, mot_remb = self.process_motion( - motion_latents, drop_motion_frames=drop_motion_frames) - - if len(mot) > 0: - x = [torch.cat([u, m], dim=1) for u, m in zip(x, mot)] - seq_lens = seq_lens + torch.tensor([r.size(1) for r in mot], - dtype=torch.long) - rope_embs = [ - torch.cat([u, m], dim=1) for u, m in zip(rope_embs, mot_remb) - ] - mask_input = [ - torch.cat([ - m, 2 * torch.ones([1, u.shape[1] - m.shape[1]], - device=m.device, - dtype=m.dtype) - ], - dim=1) for m, u in zip(mask_input, x) - ] - return x, seq_lens, rope_embs, mask_input def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0): if block_idx in self.audio_injector.injected_block_id.keys():