fix normal model loading
This commit is contained in:
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user