Handle InfiniteTalk first frame when not using the looping sampling
This commit is contained in:
+7
-5
@@ -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"
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user