From 96f7f6accddd3a080953d0744be12c6967d5b444 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 22:02:02 +0300 Subject: [PATCH] fix normal model loading --- nodes.py | 4 +- s2v/nodes.py | 6 ++- wanvideo/modules/model.py | 98 +++++++++++++++++---------------------- 3 files changed, 50 insertions(+), 58 deletions(-) diff --git a/nodes.py b/nodes.py index e6d0cc3..62f6150 100644 --- a/nodes.py +++ b/nodes.py @@ -2218,6 +2218,7 @@ class WanVideoSampler: if s2v_audio_embeds is not None: log.info(f"Using S2V audio embeddings") s2v_audio_input = s2v_audio_embeds["audio_embed_bucket"].to(device, dtype) + s2v_audio_scale = s2v_audio_embeds["audio_scale"] s2v_ref_latent = s2v_audio_embeds["ref_latent"] if s2v_ref_latent is not None: s2v_ref_latent = s2v_ref_latent.to(device, dtype) @@ -2683,7 +2684,8 @@ class WanVideoSampler: "mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling "mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs "s2v_audio_input": s2v_audio_input, # official speech-to-video audio input - "s2v_ref_latent": s2v_ref_latent # official speech-to-video reference latent + "s2v_ref_latent": s2v_ref_latent, # speech-to-video reference latent + "s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0 # speech-to-video audio scale } batch_size = 1 diff --git a/s2v/nodes.py b/s2v/nodes.py index fcd67a4..6c840b7 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -53,6 +53,7 @@ class WanVideoAddAudioEmbeds: "embeds": ("WANVIDIMAGE_EMBEDS",), "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), "frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}), + "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}) }, "optional": { "ref_latent": ("LATENT",) @@ -65,7 +66,7 @@ class WanVideoAddAudioEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, frames, audio_encoder_output, ref_latent=None): + def add(self, embeds, frames, audio_encoder_output, audio_scale, ref_latent=None): all_layers = audio_encoder_output["encoded_audio_all_layers"] audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512] @@ -98,7 +99,8 @@ class WanVideoAddAudioEmbeds: new_entry = { "audio_embed_bucket": audio_embed_bucket, "num_repeat": num_repeat, - "ref_latent": ref_latent["samples"] if ref_latent is not None else None + "ref_latent": ref_latent["samples"] if ref_latent is not None else None, + "audio_scale": audio_scale } updated = dict(embeds) updated["audio_embeds"] = new_entry diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 2b739e1..a77172b 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1591,35 +1591,37 @@ class WanModel(torch.nn.Module): self.block_mask=None #S2V + self.zero_timestep = None if cond_dim > 0: self.cond_encoder = nn.Conv3d( cond_dim, self.dim, kernel_size=self.patch_size, stride=self.patch_size) - self.enable_adain = enable_adain - self.casual_audio_encoder = CausalAudioEncoder( - dim=audio_dim, - out_dim=self.dim, - num_token=num_audio_token, - need_global=enable_adain) - all_modules, all_modules_names = torch_dfs( - self.blocks, parent_name="root.transformer_blocks") - self.audio_injector = AudioInjector_WAN( - all_modules, - all_modules_names, - dim=self.dim, - num_heads=self.num_heads, - inject_layer=audio_inject_layers, - root_net=self, - enable_adain=enable_adain, - adain_dim=self.dim, - need_adain_ont=adain_mode != "attn_norm", - ) - self.adain_mode = adain_mode - self.zero_timestep = zero_timestep + if self.model_type == 's2v': + self.enable_adain = enable_adain + self.casual_audio_encoder = CausalAudioEncoder( + dim=audio_dim, + out_dim=self.dim, + num_token=num_audio_token, + need_global=enable_adain) + all_modules, all_modules_names = torch_dfs( + self.blocks, parent_name="root.transformer_blocks") + self.audio_injector = AudioInjector_WAN( + all_modules, + all_modules_names, + dim=self.dim, + num_heads=self.num_heads, + inject_layer=audio_inject_layers, + root_net=self, + enable_adain=enable_adain, + adain_dim=self.dim, + need_adain_ont=adain_mode != "attn_norm", + ) + self.adain_mode = adain_mode + self.zero_timestep = zero_timestep - self.trainable_cond_mask = nn.Embedding(3, self.dim) + self.trainable_cond_mask = nn.Embedding(3, self.dim) @staticmethod def _prepare_blockwise_causal_attn_mask( @@ -1772,45 +1774,30 @@ class WanModel(torch.nn.Module): return hints - def audio_injector_forward(self, block_idx, hidden_states, merged_audio_emb): + def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0): if block_idx in self.audio_injector.injected_block_id.keys(): audio_attn_id = self.audio_injector.injected_block_id[block_idx] - audio_emb = merged_audio_emb # b f n c - num_frames = audio_emb.shape[1] + num_frames = audio_emb.shape[1]# b f n c - input_hidden_states = hidden_states[:, :self.original_seq_len].clone() # b (f h w) c - input_hidden_states = rearrange( - input_hidden_states, "b (t n) c -> (b t) n c", t=num_frames) + input_x = x[:, :self.original_seq_len].clone() # b (f h w) c + input_x = rearrange(input_x, "b (t n) c -> (b t) n c", t=num_frames) if self.enable_adain and self.adain_mode == "attn_norm": audio_emb_global = self.audio_emb_global - audio_emb_global = rearrange(audio_emb_global, - "b t n c -> (b t) n c") - adain_hidden_states = self.audio_injector.injector_adain_layers[ - audio_attn_id]( - input_hidden_states, temb=audio_emb_global[:, 0]) - attn_hidden_states = adain_hidden_states + audio_emb_global = rearrange(audio_emb_global,"b t n c -> (b t) n c") + attn_x = self.audio_injector.injector_adain_layers[audio_attn_id](input_x, temb=audio_emb_global[:, 0]) else: - attn_hidden_states = self.audio_injector.injector_pre_norm_feat[ - audio_attn_id]( - input_hidden_states) - audio_emb = rearrange( - audio_emb, "b t n c -> (b t) n c", t=num_frames) - attn_audio_emb = audio_emb - residual_out = self.audio_injector.injector[audio_attn_id]( - x=attn_hidden_states, - context=attn_audio_emb, - context_lens=torch.ones( - attn_hidden_states.shape[0], - dtype=torch.long, - device=attn_hidden_states.device) * attn_audio_emb.shape[1]) - residual_out = rearrange( - residual_out, "(b t) n c -> b (t n) c", t=num_frames) - hidden_states[:, :self. - original_seq_len] = hidden_states[:, :self. - original_seq_len] + residual_out + attn_x = self.audio_injector.injector_pre_norm_feat[audio_attn_id](input_x) - return hidden_states + attn_audio_emb = rearrange(audio_emb, "b t n c -> (b t) n c", t=num_frames) + residual_out = self.audio_injector.injector[audio_attn_id]( + x=attn_x , + context=attn_audio_emb * scale, + ) + residual_out = rearrange(residual_out, "(b t) n c -> b (t n) c", t=num_frames) + x[:, :self.original_seq_len].add_(residual_out) + + return x def forward( self, @@ -1857,7 +1844,8 @@ class WanModel(torch.nn.Module): mtv_freqs=None, mtv_strength=1.0, s2v_audio_input=None, - s2v_ref_latent=None + s2v_ref_latent=None, + s2v_audio_scale=1.0 ): r""" @@ -2451,7 +2439,7 @@ class WanModel(torch.nn.Module): continue x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block if self.audio_injector is not None and s2v_audio_input is not None: - x = self.audio_injector_forward(b, x, merged_audio_emb) #s2v + x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v if self.block_swap_debug: compute_end = time.perf_counter() compute_time = compute_end - compute_start