fix normal model loading

This commit is contained in:
kijai
2025-08-26 22:02:02 +03:00
parent 3c79851230
commit 96f7f6accd
3 changed files with 50 additions and 58 deletions
+3 -1
View File
@@ -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
+4 -2
View File
@@ -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
+43 -55
View File
@@ -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