Experimental: Allow HuMo to work with InfiniteTalk

Doesn't work that great, but pushing this anyway for possible future usecases
This commit is contained in:
kijai
2025-09-15 19:07:34 +03:00
parent 61c68365d2
commit 076c4a5c7e
5 changed files with 110 additions and 44 deletions
+25 -3
View File
@@ -204,7 +204,7 @@ class HuMoEmbeds:
log.info(f"HuMo set to generate {pixel_frame_num} frames")
audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
#audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
num_refs = 0
if reference_images is not None:
@@ -229,8 +229,8 @@ class HuMoEmbeds:
if reference_images is not None:
mask[:,:-num_refs] = 0
image_cond = torch.cat([zero_latents[:, :(target_shape[1]-num_refs)], samples], dim=1)
zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
#zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
#audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
else:
image_cond = zero_latents
mask = torch.zeros_like(mask)
@@ -252,14 +252,36 @@ class HuMoEmbeds:
}
return (embeds, )
class WanVideoCombineEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds_1": ("WANVIDIMAGE_EMBEDS",),
"embeds_2": ("WANVIDIMAGE_EMBEDS",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, embeds_1, embeds_2):
# Combine the two sets of embeds
combined = {**embeds_1, **embeds_2}
return (combined,)
NODE_CLASS_MAPPINGS = {
"WhisperModelLoader": WhisperModelLoader,
"HuMoEmbeds": HuMoEmbeds,
"WanVideoCombineEmbeds": WanVideoCombineEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WhisperModelLoader": "Whisper Model Loader",
"HuMoEmbeds": "HuMo Embeds",
"WanVideoCombineEmbeds": "WanVideo Combine Embeds",
}
+1 -1
View File
@@ -428,7 +428,7 @@ class WanVideoImageToVideoMultiTalk:
image_embeds = {
"multitalk_sampling": True,
"multitalk_start_image": resized_start_image if start_image is not None else None,
"num_frames": num_frames,
"frame_window_size": num_frames,
"motion_frame": motion_frame,
"target_h": H,
"target_w": W,
+64 -25
View File
@@ -2107,15 +2107,24 @@ class WanVideoSampler:
#HuMo inputs
humo_audio = image_embeds.get("humo_audio_emb", None)
if humo_audio is not None:
humo_audio = humo_audio.to(device, dtype)
humo_audio_neg = image_embeds.get("humo_audio_emb_neg", None)
humo_reference_count = image_embeds.get("humo_reference_count", 0)
num_frames = image_embeds.get("num_frames", 0)
if humo_audio is not None:
from .HuMo.nodes import get_audio_emb_window
if not multitalk_sampling:
humo_audio, _ = get_audio_emb_window(humo_audio, num_frames, frame0_idx=0)
zero_audio_pad = torch.zeros(humo_reference_count, *humo_audio.shape[1:]).to(humo_audio.device)
humo_audio = torch.cat([humo_audio, zero_audio_pad], dim=0)
humo_audio_neg = torch.zeros_like(humo_audio, dtype=humo_audio.dtype, device=humo_audio.device)
humo_audio = humo_audio.to(device, dtype)
if humo_audio_neg is not None:
humo_audio_neg = humo_audio_neg.to(device, dtype)
humo_audio_scale = image_embeds.get("humo_audio_scale", 1.0)
humo_image_cond = image_embeds.get("humo_image_cond", None)
humo_image_cond_neg = image_embeds.get("humo_image_cond_neg", None)
humo_reference_count = image_embeds.get("humo_reference_count", 0)
humo_audio_cfg_scale = image_embeds.get("humo_audio_cfg_scale", 1.0)
humo_start_percent = image_embeds.get("humo_start_percent", 0.0)
humo_end_percent = image_embeds.get("humo_end_percent", 1.0)
@@ -2178,7 +2187,7 @@ class WanVideoSampler:
}
# FantasyTalking
audio_proj = multitalk_audio_embedding = None
audio_proj = multitalk_audio_embeds = None
audio_scale = 1.0
if fantasytalking_embeds is not None:
audio_proj = fantasytalking_embeds["audio_proj"].to(device)
@@ -2191,13 +2200,13 @@ class WanVideoSampler:
# Handle single or multiple speaker embeddings
audio_features_in = multitalk_embeds.get("audio_features", None)
if audio_features_in is None:
multitalk_audio_embedding = None
multitalk_audio_embeds = None
else:
if isinstance(audio_features_in, list):
multitalk_audio_embedding = [emb.to(device, dtype) for emb in audio_features_in]
multitalk_audio_embeds = [emb.to(device, dtype) for emb in audio_features_in]
else:
# keep backward-compatibility with single tensor input
multitalk_audio_embedding = [audio_features_in.to(device, dtype)]
multitalk_audio_embeds = [audio_features_in.to(device, dtype)]
audio_scale = multitalk_embeds.get("audio_scale", 1.0)
audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0)
@@ -2205,7 +2214,7 @@ class WanVideoSampler:
if not isinstance(audio_cfg_scale, list):
audio_cfg_scale = [audio_cfg_scale] * (steps + 1)
shapes = [tuple(e.shape) for e in multitalk_audio_embedding]
shapes = [tuple(e.shape) for e in multitalk_audio_embeds]
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
# FantasyPortrait
@@ -2628,7 +2637,8 @@ class WanVideoSampler:
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, audio_proj=None, control_camera_latents=None,
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None):
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None,
humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None):
nonlocal transformer
z = z.to(dtype)
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
@@ -2746,8 +2756,8 @@ class WanVideoSampler:
else:
z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
if not multitalk_sampling and multitalk_audio_embedding is not None:
audio_embedding = multitalk_audio_embedding
if not multitalk_sampling and multitalk_audio_embeds is not None:
audio_embedding = multitalk_audio_embeds
audio_embs = []
indices = (torch.arange(4 + 1) - 2) * 1
human_num = len(audio_embedding)
@@ -2828,8 +2838,8 @@ class WanVideoSampler:
"add_cond": add_cond_input, # additional conditioning input
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk audio input
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk reference target masks
"multitalk_audio": multitalk_audio_input if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk audio input
"ref_target_masks": ref_target_masks if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk reference target masks
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
"standin_input": standin_input, # Stand-in reference input
"fantasy_portrait_input": fantasy_portrait_input, # Fantasy portrait input
@@ -2910,7 +2920,7 @@ class WanVideoSampler:
+ cfg_scale * (noise_pred_cond - noise_pred_phantom[0]))
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_phantom]
#audio cfg (fantasytalking and multitalk)
if (fantasytalking_embeds is not None or multitalk_audio_embedding is not None):
if (fantasytalking_embeds is not None or multitalk_audio_embeds is not None):
if not math.isclose(audio_cfg_scale[idx], 1.0):
if cache_state is not None and len(cache_state) != 3:
cache_state.append(None)
@@ -3396,7 +3406,8 @@ class WanVideoSampler:
text_embeds["negative_prompt_embeds"],
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose)
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose,
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,)
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
@@ -3416,7 +3427,7 @@ class WanVideoSampler:
offload = image_embeds.get("force_offload", False)
offloaded = False
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("num_frames", 81)
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
vae = image_embeds.get("vae", None)
clip_embeds = image_embeds.get("clip_context", None)
if clip_embeds is not None:
@@ -3454,9 +3465,10 @@ class WanVideoSampler:
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
audio_embedding = multitalk_audio_embedding
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
cond_frame = None
pcd_data = pcd_data_input = None
if uni3c_embeds is not None:
@@ -3564,9 +3576,11 @@ class WanVideoSampler:
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
# zero padding and vae encode for img cond
if cond_image is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if is_first_clip else cond_frame
cond_frame_num = cond_.shape[2]
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
@@ -3589,6 +3603,25 @@ class WanVideoSampler:
y = None
latent_motion_frames = noise[:, :1]
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
if humo_image_cond is not None:
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
if y is not None:
partial_humo_cond_input[:, :1] = y[:, :1]
if humo_reference_count > 0:
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
if humo_audio is not None:
if is_first_clip:
audio_embs = None
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
partial_humo_audio[-humo_reference_count:] = 0
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
timesteps.append(0.)
@@ -3723,12 +3756,14 @@ class WanVideoSampler:
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
if humo_image_cond is None or not is_first_clip:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
@@ -3761,12 +3796,15 @@ class WanVideoSampler:
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
else:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer)
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.model.clear_cache()
@@ -3823,7 +3861,7 @@ class WanVideoSampler:
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num)
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
@@ -4009,7 +4047,8 @@ class WanVideoSampler:
cfg[idx], text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input)
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input,
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg)
if bidirectional_sampling:
noise_pred_flipped, self.cache_state = predict_with_cfg(
latent_model_input_flipped,
+16 -11
View File
@@ -783,15 +783,18 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
for r in reader:
all_tensors.extend(r.tensors)
for tensor in all_tensors:
name = tensor.name
if "glob" not in name and "audio_proj" in name:
name = name.replace("audio_proj", "multitalk_audio_proj")
load_device = device
if "vace_blocks." in tensor.name:
if "vace_blocks." in name:
try:
vace_block_idx = int(tensor.name.split("vace_blocks.")[1].split(".")[0])
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
except Exception:
vace_block_idx = None
elif "blocks." in tensor.name:
elif "blocks." in name:
try:
block_idx = int(tensor.name.split("blocks.")[1].split(".")[0])
block_idx = int(name.split("blocks.")[1].split(".")[0])
except Exception:
block_idx = None
@@ -805,7 +808,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
sd[name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
sd.update(extra_sd)
del all_tensors, extra_sd
@@ -1298,16 +1301,21 @@ class WanVideoModelLoader:
class_interval=4,
attention_mode=attention_mode,
)
transformer.audio_proj = multitalk_model["proj_model"]
transformer.multitalk_audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type
extra_model_path = multitalk_model["model_path"]
extra_sd = {}
if multitalk_model_path.endswith(".gguf"):
extra_sd, extra_reader = load_gguf(extra_model_path)
extra_sd_temp, extra_reader = load_gguf(extra_model_path)
gguf_reader.append(extra_reader)
del extra_reader
else:
extra_sd = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
extra_sd_temp = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
for k, v in extra_sd_temp.items():
extra_sd[k.replace("audio_proj.", "multitalk_audio_proj.")] = v
sd.update(extra_sd)
del extra_sd
@@ -1392,9 +1400,6 @@ class WanVideoModelLoader:
from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights)
if multitalk_model is not None:
transformer.audio_proj = multitalk_model["proj_model"]
if vram_management_args is not None:
if gguf:
raise ValueError("GGUF models don't support vram management")
+4 -4
View File
@@ -2279,12 +2279,12 @@ class WanModel(torch.nn.Module):
# MultiTalk
if multitalk_audio is not None:
self.audio_proj.to(self.main_device)
self.multitalk_audio_proj.to(self.main_device)
audio_cond = multitalk_audio.to(device=x.device, dtype=x.dtype)
first_frame_audio_emb_s = audio_cond[:, :1, ...]
latter_frame_audio_emb = audio_cond[:, 1:, ...]
latter_frame_audio_emb = rearrange(latter_frame_audio_emb, "b (n_t n) w s c -> b n_t n w s c", n=4)
middle_index = self.audio_proj.seq_len // 2
middle_index = self.multitalk_audio_proj.seq_len // 2
latter_first_frame_audio_emb = latter_frame_audio_emb[:, :, :1, :middle_index+1, ...]
latter_first_frame_audio_emb = rearrange(latter_first_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c")
latter_last_frame_audio_emb = latter_frame_audio_emb[:, :, -1:, middle_index:, ...]
@@ -2292,10 +2292,10 @@ class WanModel(torch.nn.Module):
latter_middle_frame_audio_emb = latter_frame_audio_emb[:, :, 1:-1, middle_index:middle_index+1, ...]
latter_middle_frame_audio_emb = rearrange(latter_middle_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c")
latter_frame_audio_emb_s = torch.concat([latter_first_frame_audio_emb, latter_middle_frame_audio_emb, latter_last_frame_audio_emb], dim=2)
multitalk_audio_embedding = self.audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
multitalk_audio_embedding = self.multitalk_audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
human_num = len(multitalk_audio_embedding)
multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(x.dtype)
self.audio_proj.to(self.offload_device)
self.multitalk_audio_proj.to(self.offload_device)
# convert ref_target_masks to token_ref_target_masks
token_ref_target_masks = None