Context windows with FantasyTalking

This commit is contained in:
kijai
2025-04-28 20:24:41 +03:00
parent d96787b312
commit 8117c6f033
2 changed files with 26 additions and 2 deletions
+7 -2
View File
@@ -2518,6 +2518,7 @@ class WanVideoSampler:
"end_percent": unianimate_poses["end_percent"]
}
audio_proj = None
if fantasytalking_embeds is not None:
audio_proj = fantasytalking_embeds["audio_proj"].to(device)
audio_context_lens = fantasytalking_embeds["audio_context_lens"]
@@ -2759,7 +2760,7 @@ class WanVideoSampler:
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, teacache_state=None):
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, teacache_state=None):
z = z.to(dtype)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
@@ -3191,6 +3192,10 @@ class WanVideoSampler:
if has_ref:
partial_vace_context[:, 0, :, :] = vace_data[0]["context"][0][:, 0, :, :]
partial_vace_context = [partial_vace_context]
if fantasytalking_embeds is not None:
partial_audio_proj = audio_proj[:, c]
partial_latent_model_input = latent_model_input[:, c, :, :]
partial_unianim_data = None
@@ -3209,7 +3214,7 @@ class WanVideoSampler:
partial_latent_model_input,
cfg[idx], positive,
text_embeds["negative_prompt_embeds"],
timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,
timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
current_teacache)
# if callback is not None:
+19
View File
@@ -341,6 +341,25 @@ class WanT2VCrossAttention(WanSelfAttention):
# output
x = x.flatten(2)
# FantasyTalking audio attention
if audio_proj is not None:
if len(audio_proj.shape) == 4:
audio_q = q.view(b * num_latent_frames, -1, n, d) # [b, 21, l1, n, d]
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
audio_x = attention(
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode
)
audio_x = audio_x.view(b, q.size(1), n, d)
audio_x = audio_x.flatten(2)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode)
audio_x = audio_x.flatten(2)
x = x + audio_x * audio_scale
x = self.o(x)
return x