Allow T2V with HuMo

This commit is contained in:
kijai
2025-09-13 19:20:19 +03:00
parent 949647c4d0
commit 9063b99bd1
2 changed files with 40 additions and 26 deletions
+36 -23
View File
@@ -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,
+4 -3
View File
@@ -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