Handle InfiniteTalk first frame when not using the looping sampling

This commit is contained in:
kijai
2025-08-20 13:33:44 +03:00
parent ff779c9171
commit a1220f7f36
4 changed files with 33 additions and 7 deletions
+7 -5
View File
@@ -52,7 +52,8 @@ class MultiTalkModelLoader:
multitalk = {
"proj_model": multitalk_proj_model,
"sd": sd,
"is_gguf": model_path.endswith(".gguf")
"is_gguf": model_path.endswith(".gguf"),
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
}
return (multitalk,)
@@ -266,9 +267,10 @@ class WanVideoImageToVideoMultiTalk:
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
"mode": ([
"auto",
"multitalk",
"infinitetalk"
], {"default": "multitalk", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
}
}
@@ -321,7 +323,7 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MultiTalkModelLoader": "MultiTalk Model Loader",
"MultiTalkWav2VecEmbeds": "MultiTalk Wav2Vec Embeds",
"WanVideoImageToVideoMultiTalk": "WanVideo Image To Video MultiTalk"
"MultiTalkModelLoader": "Multi/InfiniteTalk Model Loader",
"MultiTalkWav2VecEmbeds": "Multi/InfiniteTalk Wav2Vec Embeds",
"WanVideoImageToVideoMultiTalk": "WanVideo Image To Video Multi/InfiniteTalk"
}
+21 -1
View File
@@ -2069,7 +2069,7 @@ class WanVideoSampler:
# extra latents (Pusa) and 5b
latents_to_insert = add_index = None
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
if (extra_latents := image_embeds.get("extra_latents", None)) is not None and transformer.multitalk_model_type.lower() != "infinitetalk":
all_indices = []
for entry in extra_latents:
add_index = entry["index"]
@@ -2725,6 +2725,15 @@ class WanVideoSampler:
latent_flipped = torch.flip(latent, dims=[1])
latent_model_input_flipped = latent_flipped.to(device)
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
latent_model_input = latent.to(device)
current_step_percentage = idx / len(timesteps)
@@ -3002,6 +3011,8 @@ class WanVideoSampler:
#region multitalk
elif multitalk_sampling:
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
offload = image_embeds.get("force_offload", False)
@@ -3399,6 +3410,15 @@ class WanVideoSampler:
**scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
if freeinit_args is not None:
current_latent = latent.clone()
+3 -1
View File
@@ -1014,6 +1014,7 @@ class WanVideoModelLoader:
if multitalk_model is not None:
if multitalk_model["is_gguf"] and not gguf:
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention
from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm
@@ -1033,8 +1034,9 @@ class WanVideoModelLoader:
attention_mode=attention_mode,
)
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
log.info("MultiTalk model detected, patching model...")
log.info(f"{multitalk_model_type} detected, patching model...")
transformer.audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type
sd.update(multitalk_model["sd"])
# Additional cond latents
+2
View File
@@ -1260,6 +1260,8 @@ class WanModel(torch.nn.Module):
self.video_attention_split_steps = []
self.lora_scheduling_enabled = False
self.multitalk_model_type = None
# embeddings
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)