From 076c4a5c7ed0d213c2c6b95c49d7ce37ae00fbc4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 15 Sep 2025 19:07:34 +0300 Subject: [PATCH] Experimental: Allow HuMo to work with InfiniteTalk Doesn't work that great, but pushing this anyway for possible future usecases --- HuMo/nodes.py | 28 ++++++++++-- multitalk/nodes.py | 2 +- nodes.py | 89 ++++++++++++++++++++++++++++----------- nodes_model_loading.py | 27 +++++++----- wanvideo/modules/model.py | 8 ++-- 5 files changed, 110 insertions(+), 44 deletions(-) diff --git a/HuMo/nodes.py b/HuMo/nodes.py index 1bb15a5..9659d63 100644 --- a/HuMo/nodes.py +++ b/HuMo/nodes.py @@ -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", } diff --git a/multitalk/nodes.py b/multitalk/nodes.py index 4e5690b..fed8fb8 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -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, diff --git a/nodes.py b/nodes.py index aaa904a..6636285 100644 --- a/nodes.py +++ b/nodes.py @@ -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, diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 4fe04f7..ccae150 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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") diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 68d1c3e..b27e043 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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