diff --git a/HuMo/nodes.py b/HuMo/nodes.py index 1a56a22..f6f3d89 100644 --- a/HuMo/nodes.py +++ b/HuMo/nodes.py @@ -5,7 +5,7 @@ import os import json import torchaudio -from comfy.utils import load_torch_file +from comfy.utils import load_torch_file, common_upscale import comfy.model_management as mm from accelerate import init_empty_weights @@ -120,16 +120,18 @@ class HuMoEmbeds: @classmethod def INPUT_TYPES(s): return {"required": { - "whisper_model": ("WHISPERMODEL",), - "vae": ("WANVAE", ), "num_frames": ("INT", {"default": 81, "min": -1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}), - "reference_images": ("IMAGE", {"tooltip": "reference images for the humo model"}), + "width": ("INT", {"default": 832, "min": 64, "max": 4096, "step": 16}), + "height": ("INT", {"default": 480, "min": 64, "max": 4096, "step": 16}), "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}), "audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}), "audio_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to start applying audio conditioning"}), "audio_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to stop applying audio conditioning"}) }, "optional" : { + "whisper_model": ("WHISPERMODEL",), + "vae": ("WANVAE", ), + "reference_images": ("IMAGE", {"tooltip": "reference images for the humo model"}), "audio": ("AUDIO",), } } @@ -139,7 +141,11 @@ class HuMoEmbeds: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, whisper_model, vae, reference_images, num_frames, audio_scale, audio_cfg_scale, audio_start_percent, audio_end_percent, audio=None): + def process(self, num_frames, width, height, audio_scale, audio_cfg_scale, audio_start_percent, audio_end_percent, whisper_model=None, vae=None, reference_images=None, audio=None): + if reference_images is not None and vae is None: + raise ValueError("VAE is required when reference images are provided") + if whisper_model is None and audio is not None: + raise ValueError("Whisper model is required when audio is provided") model = whisper_model["model"] feature_extractor = whisper_model["feature_extractor"] dtype = whisper_model["dtype"] @@ -199,38 +205,45 @@ class HuMoEmbeds: audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0) - samples, = WanVideoEncodeLatentBatch.encode(self, vae, reference_images, False, 0, 0, 0, 0) - samples = samples["samples"].transpose(0, 2).squeeze(0) - - C, T, H, W = samples.shape - - target_shape = (16, latent_frame_num + T, - H * 8 // 8, - W * 8 // 8) + num_refs = 0 + if reference_images is not None: + if reference_images.shape[1] != height or reference_images.shape[2] != width: + reference_images_in = common_upscale(reference_images.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(-1, 1) + else: + reference_images_in = reference_images + samples, = WanVideoEncodeLatentBatch.encode(self, vae, reference_images_in, False, 0, 0, 0, 0) + samples = samples["samples"].transpose(0, 2).squeeze(0) + num_refs = samples.shape[1] vae.to(device) - zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*T, H * 8, W * 8, device=device, dtype=vae.dtype) - zero_latents = vae.encode(zero_frames, device=device)[0].to(samples.device) + zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*num_refs, height, width, device=device, dtype=vae.dtype) + zero_latents = vae.encode(zero_frames, device=device)[0].to(offload_device) vae.model.clear_cache() vae.to(offload_device) mm.soft_empty_cache() - mask = torch.ones(4, target_shape[1], target_shape[2], target_shape[3], device=samples.device, dtype=vae.dtype) - mask[:,:-T] = 0 - image_cond = torch.cat([zero_latents[:, :(target_shape[1]-T)], samples], dim=1) + target_shape = (16, latent_frame_num + num_refs, height // 8, width // 8) + + mask = torch.ones(4, target_shape[1], target_shape[2], target_shape[3], device=offload_device, dtype=vae.dtype) + 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) + else: + image_cond = zero_latents + mask = torch.zeros_like(mask) image_cond = torch.cat([mask, image_cond], dim=0) image_cond_neg = torch.cat([mask, zero_latents], dim=0) - zero_audio_pad = torch.zeros(T, *audio_emb.shape[1:]).to(audio_emb.device) - audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0) - audio_emb_neg = torch.zeros_like(audio_emb, dtype=audio_emb.dtype, device=audio_emb.device) + embeds = { "humo_audio_emb": audio_emb, - "humo_audio_emb_neg": audio_emb_neg, + "humo_audio_emb_neg": torch.zeros_like(audio_emb, dtype=audio_emb.dtype, device=audio_emb.device), "humo_image_cond": image_cond, "humo_image_cond_neg": image_cond_neg, - "humo_reference_count": T, + "humo_reference_count": num_refs, "target_shape": target_shape, "num_frames": pixel_frame_num, "humo_audio_scale": audio_scale, diff --git a/nodes.py b/nodes.py index 556587a..b7d0c00 100644 --- a/nodes.py +++ b/nodes.py @@ -2676,7 +2676,8 @@ class WanVideoSampler: elif humo_image_cond is not None: if context_window is not None: image_cond_input = humo_image_cond[:, context_window].to(z) - image_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:] + if humo_reference_count > 0: + image_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:] else: image_cond_input = humo_image_cond.to(z) elif image_cond is not None: @@ -4085,7 +4086,7 @@ class WanVideoSampler: callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach() #elif phantom_latents is not None: # callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach() - elif humo_image_cond is not None: + elif humo_image_cond is not None and humo_reference_count > 0: callback_latent = (latent_model_input[:,:-humo_reference_count].to(device) - noise_pred[:,:-humo_reference_count].to(device) * t.to(device) / 1000).detach() else: callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach() @@ -4107,7 +4108,7 @@ class WanVideoSampler: if phantom_latents is not None: latent = latent[:,:-phantom_latents.shape[1]] - if humo_image_cond is not None: + if humo_image_cond is not None and humo_reference_count > 0: latent = latent[:,:-humo_reference_count] cache_states = None