Squashed commit of the following:

commit fda0fe6e0c21eb10276ae302cd88b6cbcf5b36b5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:55:00 2025 +0300

    Create wanvideo_HuMo_example_01.json

commit cffe3039c3d2fbacd4803329bf31b5fdc45215ba
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:30:49 2025 +0300

    Update model.py

commit ddce018a5a6ffeb926860342889b12efb7343ec0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:29:27 2025 +0300

    cleanup

commit 8c021b8b3f66144804e500e74aa6f2be52f9f9fc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:23:27 2025 +0300

    avoid compile graph break

commit ef9c7732042261581b4bba6d980def78633f56dc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:16:13 2025 +0300

    Allow using whisper model without decoder layers

commit 8d0ba29ee84d14be6084ecbcee1cbc9414128fb3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 15:55:26 2025 +0300

    start/end percent for HuMo audio

commit bfe0d358a8820240f262351e61cfb979cb9a47ff
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 15:37:11 2025 +0300

    cleanup

commit e563ae317f24a7f5751b43cdef4bffbaeaea5114
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 14:02:21 2025 +0300

    Make audio work

commit 95855196c51b1124a19079b746c5d12ec70d9026
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 18:10:04 2025 +0300

    cfg

commit d5a18b090fe719b7f0b00a0f68e6824685598313
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 03:10:15 2025 +0300

    wrong way around

commit 34c8c4842c14002fe4694dfa23c24a65b7ea39d0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 03:01:45 2025 +0300

    Update nodes.py

commit 47d1e2ab5f3e1483782d0739f5b51cdd33707c36
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 02:47:03 2025 +0300

    update

    image inputs are working but audio still doesn't do anything

commit 67890d816a64459944091cb01478c1e0ec4c4a82
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 21:13:12 2025 +0300

    update

commit dbcef53405bb78feae4c5d2c6b310b76e4ef9949
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 17:09:37 2025 +0300

    Update model.py

commit 92c9aac51f4d37988757510a4e57179834cc5de2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 16:15:39 2025 +0300

    init

    untested as no weights released as of yet
This commit is contained in:
kijai
2025-09-13 16:55:28 +03:00
parent 8ce4432fef
commit 38fd791a77
8 changed files with 2368 additions and 110 deletions
+87
View File
@@ -0,0 +1,87 @@
import torch
from einops import rearrange
from torch import nn
from einops import rearrange
class WanRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
r"""
Args:
x(Tensor): Shape [B, L, C]
"""
return self._norm(x.float()).type_as(x) * self.weight
def _norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
class DummyAdapterLayer(nn.Module):
def __init__(self, layer):
super().__init__()
self.layer = layer
def forward(self, *args, **kwargs):
return self.layer(*args, **kwargs)
class AudioProjModel(nn.Module):
def __init__(
self,
seq_len=5,
blocks=13, # add a new parameter blocks
channels=768, # add a new parameter channels
intermediate_dim=512,
output_dim=1536,
context_tokens=16,
):
super().__init__()
self.seq_len = seq_len
self.blocks = blocks
self.channels = channels
self.input_dim = seq_len * blocks * channels # update input_dim to be the product of blocks and channels.
self.intermediate_dim = intermediate_dim
self.context_tokens = context_tokens
self.output_dim = output_dim
# define multiple linear layers
self.audio_proj_glob_1 = DummyAdapterLayer(nn.Linear(self.input_dim, intermediate_dim))
self.audio_proj_glob_2 = DummyAdapterLayer(nn.Linear(intermediate_dim, intermediate_dim))
self.audio_proj_glob_3 = DummyAdapterLayer(nn.Linear(intermediate_dim, context_tokens * output_dim))
self.audio_proj_glob_norm = DummyAdapterLayer(nn.LayerNorm(output_dim))
self.initialize_weights()
def initialize_weights(self):
# Initialize transformer layers:
def _basic_init(module):
if isinstance(module, nn.Linear):
torch.nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
self.apply(_basic_init)
def forward(self, audio_embeds):
video_length = audio_embeds.shape[1]
audio_embeds = rearrange(audio_embeds, "bz f w b c -> (bz f) w b c")
batch_size, window_size, blocks, channels = audio_embeds.shape
audio_embeds = audio_embeds.view(batch_size, window_size * blocks * channels)
audio_embeds = torch.relu(self.audio_proj_glob_1(audio_embeds))
audio_embeds = torch.relu(self.audio_proj_glob_2(audio_embeds))
context_tokens = self.audio_proj_glob_3(audio_embeds).reshape(batch_size, self.context_tokens, self.output_dim)
context_tokens = self.audio_proj_glob_norm(context_tokens)
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
return context_tokens
+249
View File
@@ -0,0 +1,249 @@
import folder_paths
import torch
import torch.nn.functional as F
import os
import json
import torchaudio
from comfy.utils import load_torch_file
import comfy.model_management as mm
from accelerate import init_empty_weights
from ..utils import set_module_tensor_to_device
from ..nodes import WanVideoEncodeLatentBatch
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def linear_interpolation_fps(features, input_fps, output_fps, output_len=None):
features = features.transpose(1, 2) # [1, C, T]
seq_len = features.shape[2] / float(input_fps)
if output_len is None:
output_len = int(seq_len * output_fps)
output_features = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return output_features.transpose(1, 2)
def get_audio_emb_window(audio_emb, frame_num, frame0_idx, audio_shift=2):
zero_audio_embed = torch.zeros((audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
zero_audio_embed_3 = torch.zeros((3, audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
iter_ = 1 + (frame_num - 1) // 4
audio_emb_wind = []
for lt_i in range(iter_):
if lt_i == 0:
st = frame0_idx + lt_i - 2
ed = frame0_idx + lt_i + 3
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
wind_feat = torch.cat((zero_audio_embed_3, wind_feat), dim=0)
else:
st = frame0_idx + 1 + 4 * (lt_i - 1) - audio_shift
ed = frame0_idx + 1 + 4 * lt_i + audio_shift
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
audio_emb_wind.append(wind_feat)
audio_emb_wind = torch.stack(audio_emb_wind, dim=0)
return audio_emb_wind, ed - audio_shift
class WhisperModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("audio_encoders"), {"tooltip": "These models are loaded from the 'ComfyUI/models/wav2vec2' or 'ComfyUI/models/audio_encoders' folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
},
}
RETURN_TYPES = ("WHISPERMODEL",)
RETURN_NAMES = ("whisper_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device):
from transformers import WhisperConfig, WhisperModel, WhisperFeatureExtractor
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
if load_device == "offload_device":
transformer_load_device = offload_device
else:
transformer_load_device = device
config_path = os.path.join(script_directory, "whisper_config.json")
whisper_config = WhisperConfig(**json.load(open(config_path)))
with init_empty_weights():
whisper = WhisperModel(whisper_config).eval()
whisper.decoder = None # we only need the encoder
feature_extractor_config = {
"chunk_length": 30,
"feature_extractor_type": "WhisperFeatureExtractor",
"feature_size": 128,
"hop_length": 160,
"n_fft": 400,
"n_samples": 480000,
"nb_max_frames": 3000,
"padding_side": "right",
"padding_value": 0.0,
"processor_class": "WhisperProcessor",
"return_attention_mask": False,
"sampling_rate": 16000
}
feature_extractor = WhisperFeatureExtractor(**feature_extractor_config)
model_path = folder_paths.get_full_path_or_raise("audio_encoders", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
for name, param in whisper.named_parameters():
key = "model." + name
value=sd[key]
set_module_tensor_to_device(whisper, name, device=offload_device, dtype=base_dtype, value=value)
whisper_model = {
"feature_extractor": feature_extractor,
"model": whisper,
"dtype": base_dtype,
}
return (whisper_model,)
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"}),
"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" : {
"audio": ("AUDIO",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds", )
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):
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sampling_rate = 16000
if audio is not None:
audio_input = audio["waveform"][0]
sample_rate = audio["sample_rate"]
if sample_rate != sampling_rate:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sampling_rate)
if audio_input.shape[1] == 2:
audio_input = audio_input.mean(dim=0, keepdim=False)
else:
audio_input = audio_input[0]
model.to(device)
audio_len = len(audio_input) // 640
# feature extraction
audio_features = []
window = 750*640
for i in range(0, len(audio_input), window):
audio_feature = feature_extractor(audio_input[i:i+window], sampling_rate=sampling_rate, return_tensors="pt").input_features
audio_features.append(audio_feature)
audio_features = torch.cat(audio_features, dim=-1).to(device, dtype)
# preprocess
window = 3000
audio_prompts = []
for i in range(0, audio_features.shape[-1], window):
audio_prompt = model.encoder(audio_features[:,:,i:i+window], output_hidden_states=True).hidden_states
audio_prompt = torch.stack(audio_prompt, dim=2)
audio_prompts.append(audio_prompt)
model.to(offload_device)
audio_prompts = torch.cat(audio_prompts, dim=1)
audio_prompts = audio_prompts[:,:audio_len*2]
feat0 = linear_interpolation_fps(audio_prompts[:, :, 0: 8].mean(dim=2), 50, 25)
feat1 = linear_interpolation_fps(audio_prompts[:, :, 8: 16].mean(dim=2), 50, 25)
feat2 = linear_interpolation_fps(audio_prompts[:, :, 16: 24].mean(dim=2), 50, 25)
feat3 = linear_interpolation_fps(audio_prompts[:, :, 24: 32].mean(dim=2), 50, 25)
feat4 = linear_interpolation_fps(audio_prompts[:, :, 32], 50, 25)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
else:
audio_emb = torch.zeros(num_frames, 5, 1280, device=device)
audio_len = num_frames
frame_num = num_frames if num_frames != -1 else audio_len
frame_num = 4 * ((frame_num - 1) // 4) + 1
audio_emb, _ = get_audio_emb_window(audio_emb, 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, (num_frames - 1) // 4 + 1 + T,
H * 8 // 8,
W * 8 // 8)
vae.to(device)
zero_frames = torch.zeros(1, 3, num_frames + 4*T, H * 8, W * 8, device=device, dtype=vae.dtype)
zero_latents = vae.encode(zero_frames, device=device)[0].to(samples.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)
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_image_cond": image_cond,
"humo_image_cond_neg": image_cond_neg,
"humo_reference_count": T,
"target_shape": target_shape,
"num_frames": num_frames,
"humo_audio_scale": audio_scale,
"humo_audio_cfg_scale": audio_cfg_scale,
"humo_start_percent": audio_start_percent,
"humo_end_percent": audio_end_percent,
}
return (embeds, )
NODE_CLASS_MAPPINGS = {
"WhisperModelLoader": WhisperModelLoader,
"HuMoEmbeds": HuMoEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WhisperModelLoader": "Whisper Model Loader",
"HuMoEmbeds": "HuMo Embeds",
}
+50
View File
@@ -0,0 +1,50 @@
{
"_name_or_path": "openai/whisper-large-v3",
"activation_dropout": 0.0,
"activation_function": "gelu",
"apply_spec_augment": false,
"architectures": [
"WhisperForConditionalGeneration"
],
"attention_dropout": 0.0,
"begin_suppress_tokens": [
220,
50257
],
"bos_token_id": 50257,
"classifier_proj_size": 256,
"d_model": 1280,
"decoder_attention_heads": 20,
"decoder_ffn_dim": 5120,
"decoder_layerdrop": 0.0,
"decoder_layers": 32,
"decoder_start_token_id": 50258,
"dropout": 0.0,
"encoder_attention_heads": 20,
"encoder_ffn_dim": 5120,
"encoder_layerdrop": 0.0,
"encoder_layers": 32,
"eos_token_id": 50257,
"init_std": 0.02,
"is_encoder_decoder": true,
"mask_feature_length": 10,
"mask_feature_min_masks": 0,
"mask_feature_prob": 0.0,
"mask_time_length": 10,
"mask_time_min_masks": 2,
"mask_time_prob": 0.05,
"max_length": 448,
"max_source_positions": 1500,
"max_target_positions": 448,
"median_filter_width": 7,
"model_type": "whisper",
"num_hidden_layers": 32,
"num_mel_bins": 128,
"pad_token_id": 50256,
"scale_embedding": false,
"torch_dtype": "float16",
"transformers_version": "4.36.0.dev0",
"use_cache": true,
"use_weighted_layer_sum": false,
"vocab_size": 51866
}
+9
View File
@@ -43,6 +43,13 @@ except Exception as e:
MTV_NODE_CLASS_MAPPINGS = {}
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .HuMo.nodes import NODE_CLASS_MAPPINGS as HUMO_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as HUMO_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
print(f"HuMo nodes not available due to error in importing them: {e}")
HUMO_NODE_CLASS_MAPPINGS = {}
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
@@ -60,6 +67,7 @@ NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -78,5 +86,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
File diff suppressed because it is too large Load Diff
+101 -74
View File
@@ -1917,6 +1917,7 @@ class WanVideoSampler:
fun_or_fl2v_model = has_ref = drop_last = False
phantom_latents = fun_ref_image = ATI_tracks = None
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
humo_audio = humo_audio_neg = None
#I2V
image_cond = image_embeds.get("image_embeds", None)
@@ -2103,6 +2104,24 @@ class WanVideoSampler:
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
#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)
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)
if not isinstance(humo_audio_cfg_scale, list):
humo_audio_cfg_scale = [humo_audio_cfg_scale] * (steps + 1)
latent_video_length = noise.shape[1]
# Initialize FreeInit filter if enabled
@@ -2654,6 +2673,12 @@ class WanVideoSampler:
elif ATI_tracks is not None and ((ati_start_percent <= current_step_percentage <= ati_end_percent) or
(ati_end_percent > 0 and idx == 0 and current_step_percentage >= ati_start_percent)):
image_cond_input = image_cond_ati.to(z)
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:]
else:
image_cond_input = humo_image_cond.to(z)
elif image_cond is not None:
if reverse_time: # Flip the image condition
image_cond_input = torch.cat([
@@ -2759,12 +2784,25 @@ class WanVideoSampler:
(s2v_pose_end_percent > 0 and idx == 0 and current_step_percentage >= s2v_pose_start_percent)):
s2v_pose = None
if humo_audio is not None and ((humo_start_percent <= current_step_percentage <= humo_end_percent) or \
(humo_end_percent > 0 and idx == 0 and current_step_percentage >= humo_start_percent)):
humo_audio_input = humo_audio
humo_audio_input_neg = humo_audio_neg if humo_audio_neg is not None else None
else:
humo_audio_input = humo_audio_input_neg = None
base_params = {
'x': [z], # latent
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
'clip_fea': clip_fea, # clip features
'seq_len': seq_len, # sequence length
'device': device, # main device
'freqs': freqs, # rope freqs
't': timestep, # current timestep
'is_uncond': False, # is unconditional
'current_step': idx, # current step
'current_step_percentage': current_step_percentage, # current step percentage
'last_step': len(timesteps) - 1 == idx, # is last step
'control_lora_enabled': control_lora_enabled, # control lora toggle for patch embed selection
'enhance_enabled': enhance_enabled, # enhance-a-video toggle
@@ -2796,7 +2834,9 @@ class WanVideoSampler:
"s2v_ref_motion": s2v_ref_motion, # speech-to-video reference motion latent
"s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0, # speech-to-video audio scale
"s2v_pose": s2v_pose if s2v_pose is not None else None, # speech-to-video pose control
"s2v_motion_frames": s2v_motion_frames, # speech-to-video motion frames
"s2v_motion_frames": s2v_motion_frames, # speech-to-video motion frames,
"humo_audio": humo_audio_input, # humo audio input
"humo_audio_scale": humo_audio_scale if humo_audio is not None else 1.0, # humo audio scale
}
batch_size = 1
@@ -2809,98 +2849,88 @@ class WanVideoSampler:
try:
if not batched_cfg:
#cond
#conditional (positive) pass
noise_pred_cond, cache_state_cond = transformer(
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
context=positive_embeds,
pred_id=cache_state[0] if cache_state else None,
vace_data=vace_data, attn_cond=attn_cond,
**base_params
)
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
noise_pred_cond = noise_pred_cond[0]
if math.isclose(cfg_scale, 1.0):
if use_fresca:
noise_pred_cond = fourier_filter(
noise_pred_cond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
noise_pred_cond = fourier_filter(noise_pred_cond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
return noise_pred_cond, [cache_state_cond]
#uncond
if fantasytalking_embeds is not None:
if not math.isclose(audio_cfg_scale[idx], 1.0):
base_params['audio_proj'] = None
#unconditional (negative) pass
base_params['is_uncond'] = True
base_params['clip_fea'] = clip_fea_neg if clip_fea_neg is not None else clip_fea
if humo_audio_input_neg is not None:
base_params['humo_audio'] = humo_audio_input_neg
noise_pred_uncond, cache_state_uncond = transformer(
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
y=[image_cond_input] if image_cond_input is not None else None,
is_uncond=True, current_step_percentage=current_step_percentage,
context=negative_embeds if humo_audio_input_neg is None else positive_embeds, #ti
pred_id=cache_state[1] if cache_state else None,
vace_data=vace_data, attn_cond=attn_cond_neg,
**base_params
)
noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device)
**base_params)
noise_pred_uncond = noise_pred_uncond[0]
# HuMo
if humo_audio_input_neg is not None and not math.isclose(humo_audio_cfg_scale[idx], 1.0):
if len(cache_state) !=3:
cache_state.append(None)
if t > 980 and humo_image_cond_neg is not None: # use image cond for first timesteps
base_params['y'] = [humo_image_cond_neg.to(z)]
noise_pred_humo_audio_uncond, cache_state_humo = transformer(
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + humo_audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_humo_audio_uncond[0])
+ (cfg_scale - 2.0) * (noise_pred_humo_audio_uncond[0] - noise_pred_uncond))
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_humo]
#phantom
if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0):
noise_pred_phantom, cache_state_phantom = transformer(
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
y=[image_cond_input] if image_cond_input is not None else None,
is_uncond=True, current_step_percentage=current_step_percentage,
pred_id=cache_state[2] if cache_state else None,
vace_data=None,
**base_params
)
noise_pred_phantom = noise_pred_phantom[0].to(intermediate_device)
noise_pred = noise_pred_uncond + phantom_cfg_scale[idx] * (noise_pred_phantom - noise_pred_uncond) + cfg_scale * (noise_pred_cond - noise_pred_phantom)
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + phantom_cfg_scale[idx] * (noise_pred_phantom[0] - noise_pred_uncond)
+ cfg_scale * (noise_pred_cond - noise_pred_phantom[0]))
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_phantom]
#fantasytalking
if fantasytalking_embeds is not None:
#audio cfg (fantasytalking and multitalk)
if (fantasytalking_embeds is not None or multitalk_audio_embedding 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)
base_params['audio_proj'] = None
# Set audio parameters to None/zeros based on type
if fantasytalking_embeds is not None:
base_params['audio_proj'] = None
audio_context = positive_embeds
else: # multitalk
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
audio_context = negative_embeds
noise_pred_no_audio, cache_state_audio = transformer(
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
context=audio_context, is_uncond=False,
pred_id=cache_state[2] if cache_state else None,
vace_data=vace_data,
**base_params
)
noise_pred_no_audio = noise_pred_no_audio[0].to(intermediate_device)
noise_pred = (
noise_pred_uncond
+ cfg_scale * (noise_pred_no_audio - noise_pred_uncond)
+ audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_no_audio)
)
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_audio]
elif multitalk_audio_embedding 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)
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
noise_pred_no_audio, cache_state_audio = transformer(
[z], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
pred_id=cache_state[2] if cache_state else None,
vace_data=vace_data,
**base_params
)
noise_pred_no_audio = noise_pred_no_audio[0].to(intermediate_device)
noise_pred = (
noise_pred_no_audio
+ cfg_scale * (noise_pred_cond - noise_pred_uncond)
+ audio_cfg_scale[idx] * (noise_pred_uncond - noise_pred_no_audio)
)
**base_params)
noise_pred = (noise_pred_uncond
+ cfg_scale * (noise_pred_no_audio[0] - noise_pred_uncond)
+ audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_no_audio[0]))
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_audio]
#batched
else:
base_params['z'] = [z] * 2
base_params['y'] = [image_cond_input] * 2 if image_cond_input is not None else None
base_params['clip_fea'] = torch.cat([clip_fea, clip_fea], dim=0)
cache_state_uncond = None
[noise_pred_cond, noise_pred_uncond], cache_state_cond = transformer(
[z] + [z], context=positive_embeds + negative_embeds,
y=[image_cond_input] + [image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea.repeat(2,1,1), is_uncond=False, current_step_percentage=current_step_percentage,
context=positive_embeds + negative_embeds, is_uncond=False,
pred_id=cache_state[0] if cache_state else None,
**base_params
)
@@ -2919,7 +2949,6 @@ class WanVideoSampler:
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
noise_pred_uncond_scaled = noise_pred_uncond * alpha
if use_tangential:
@@ -2932,18 +2961,12 @@ class WanVideoSampler:
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(
noise_pred_cond - noise_pred_uncond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
filtered_cond = fourier_filter(noise_pred_cond - noise_pred_uncond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
return noise_pred, [cache_state_cond, cache_state_uncond]
if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb
@@ -4055,6 +4078,8 @@ 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:
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()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
@@ -4075,6 +4100,8 @@ class WanVideoSampler:
if phantom_latents is not None:
latent = latent[:,:-phantom_latents.shape[1]]
if humo_image_cond is not None:
latent = latent[:,:-humo_reference_count]
cache_states = None
if cache_args is not None:
+4 -1
View File
@@ -761,7 +761,7 @@ class WanVideoSetLoRAs:
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer"}
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob"}
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
cnt = 0
@@ -1112,6 +1112,8 @@ class WanVideoModelLoader:
ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0]
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
is_humo = "audio_proj.audio_proj_glob_1.layer.weight" in sd
model_type = "t2v"
if "audio_injector.injector.0.k.weight" in sd:
model_type = "s2v"
@@ -1226,6 +1228,7 @@ class WanVideoModelLoader:
"enable_adain": True if "audio_injector.injector_adain_layers.0.linear.weight" in sd else False,
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0,
"zero_timestep": model_type == "s2v",
"humo_audio": is_humo,
}
+97 -35
View File
@@ -14,7 +14,6 @@ except:
from .attention import attention
import numpy as np
from copy import deepcopy
from tqdm import tqdm
import gc
@@ -25,10 +24,7 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from ...MTV.mtv import apply_rotary_emb
#from .s2v.motioner import MotionerTransformers, FramePackMotioner, rope_precompute
#from comfy.ldm.wan.model import FramePackMotioner
class FramePackMotioner(nn.Module):
class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model
def __init__(
self,
inner_dim=1024,
@@ -362,7 +358,8 @@ class WanSelfAttention(nn.Module):
num_heads,
qk_norm=True,
eps=1e-6,
attention_mode='sdpa'):
attention_mode='sdpa',
kv_dim=None):
assert out_features % num_heads == 0
super().__init__()
self.dim = out_features
@@ -379,8 +376,12 @@ class WanSelfAttention(nn.Module):
# layers
self.q = nn.Linear(in_features, out_features)
self.k = nn.Linear(in_features, out_features)
self.v = nn.Linear(in_features, out_features)
if kv_dim is not None:
self.k = nn.Linear(kv_dim, out_features)
self.v = nn.Linear(kv_dim, out_features)
else:
self.k = nn.Linear(in_features, out_features)
self.v = nn.Linear(in_features, out_features)
self.o = nn.Linear(in_features, out_features)
self.norm_q = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
@@ -587,8 +588,8 @@ class LoRALinearLayer(nn.Module):
#region crossattn
class WanT2VCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps)
def __init__(self, in_features, out_features, num_heads, kv_dim=None, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim)
self.attention_mode = attention_mode
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
@@ -721,6 +722,45 @@ class WanI2VCrossAttention(WanSelfAttention):
x = x + adapter_x * ip_scale
return self.o(x)
class WanHuMoCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, kv_dim=None, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim)
self.attention_mode = attention_mode
def forward(self, x, context, grid_sizes, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
q = self.norm_q(self.q(x)).view(b, -1, n, d)
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
# Handle video spatial structure
hlen_wlen = grid_sizes[0][1] * grid_sizes[0][2]
q = q.reshape(-1, hlen_wlen, n, d)
# Handle audio temporal structure (16 tokens per frame)
k = k.reshape(-1, 16, n, d)
v = v.reshape(-1, 16, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode)
x_text = x_text.view(b, -1, n, d).flatten(2)
x = x_text
return self.o(x)
class AudioCrossAttentionWrapper(nn.Module):
def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, kv_dim=None):
super().__init__()
self.audio_cross_attn = WanHuMoCrossAttention(in_features, out_features, num_heads, kv_dim=kv_dim)
self.norm1_audio = WanLayerNorm(out_features, eps, elementwise_affine=True)
def forward(self, x, audio, grid_sizes, humo_audio_scale=1.0):
x = x + self.audio_cross_attn(self.norm1_audio(x), audio, grid_sizes) * humo_audio_scale
return x
class MTVCrafterMotionAttention(WanSelfAttention):
@@ -768,7 +808,8 @@ class WanAttentionBlock(nn.Module):
eps=1e-6,
attention_mode='sdpa',
rope_func="comfy",
use_motion_attn=False
use_motion_attn=False,
use_humo_audio_attn=False,
):
super().__init__()
self.dim = out_features
@@ -813,6 +854,10 @@ class WanAttentionBlock(nn.Module):
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
self.seg_idx = None
# HuMo audio cross-attn
if use_humo_audio_attn:
self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536)
#@torch.compiler.disable()
def get_mod(self, e):
if e.dim() == 3:
@@ -877,26 +922,21 @@ class WanAttentionBlock(nn.Module):
num_latent_frames=21,
original_seq_len=None,
enhance_enabled=False,
block_mask=None,
nag_params={},
nag_context=None,
is_uncond=False,
multitalk_audio_embedding=None,
ref_target_masks=None,
human_num=0,
inner_t=None,
inner_c=None,
inner_t=None, inner_c=None,
cross_freqs=None,
x_ip=None,
e_ip=None,
x_ip=None, e_ip=None,
freqs_ip=None,
adapter_proj=None,
ip_scale=1.0,
reverse_time=False,
mtv_motion_tokens=None,
mtv_motion_rotary_emb=None,
mtv_strength=1.0,
mtv_freqs=None
mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None,
humo_audio_input=None, humo_audio_scale=1.0,
):
r"""
Args:
@@ -1054,7 +1094,9 @@ class WanAttentionBlock(nn.Module):
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale,
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength)
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength,
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale
)
else:
if self.rope_func == "comfy_chunked":
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
@@ -1074,7 +1116,8 @@ class WanAttentionBlock(nn.Module):
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
audio_proj, audio_scale, num_latent_frames, nag_params,
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength):
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength,
humo_audio_input, humo_audio_scale):
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
audio_proj=audio_proj, audio_scale=audio_scale,
@@ -1092,6 +1135,10 @@ class WanAttentionBlock(nn.Module):
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
x = x + x_motion * mtv_strength
# HuMo Audio Cross-Attention
if humo_audio_input is not None:
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
if self.rope_func == "comfy_chunked" and not self.zero_timestep:
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
else:
@@ -1415,7 +1462,8 @@ class WanModel(torch.nn.Module):
enable_adain=False,
adain_mode="attn_norm",
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
zero_timestep=False
zero_timestep=False,
humo_audio=False,
):
r"""
Initialize the diffusion model backbone.
@@ -1525,6 +1573,8 @@ class WanModel(torch.nn.Module):
self.multitalk_model_type = "none"
self.humo_audio = humo_audio
# embeddings
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
@@ -1577,7 +1627,7 @@ class WanModel(torch.nn.Module):
self.blocks = nn.ModuleList([
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn))
attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio)
for i in range(num_layers)
])
#MTV Crafter
@@ -1620,8 +1670,6 @@ class WanModel(torch.nn.Module):
else:
self.control_adapter = None
self.block_mask=None
#S2V
self.zero_timestep = self.audio_injector = self.trainable_cond_mask =None
if cond_dim > 0:
@@ -1661,6 +1709,12 @@ class WanModel(torch.nn.Module):
self.adain_mode = adain_mode
self.zero_timestep = zero_timestep
# HuMo Audio
if self.humo_audio:
from ...HuMo.audio_proj import AudioProjModel
self.audio_proj = AudioProjModel(seq_len=8, blocks=5, channels=1280,
intermediate_dim=512, output_dim=1536, context_tokens=16)
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
# Clamp blocks_to_swap to valid range
@@ -1893,6 +1947,8 @@ class WanModel(torch.nn.Module):
s2v_ref_motion=None,
s2v_pose=None,
s2v_motion_frames=[1, 0],
humo_audio=None,
humo_audio_scale=1.0,
):
r"""
@@ -2251,6 +2307,16 @@ class WanModel(torch.nn.Module):
token_ref_target_masks = token_ref_target_masks.view(token_ref_target_masks.shape[0], -1)
token_ref_target_masks = token_ref_target_masks.to(x.dtype).to(device)
humo_audio_input = None
if humo_audio is not None:
humo_audio_input = self.audio_proj(humo_audio.unsqueeze(0)).permute(0, 3, 1, 2)
humo_audio_seq_len = torch.tensor(humo_audio.shape[2] * humo_audio_input.shape[3], device=device)
humo_audio_input = humo_audio_input.flatten(2).transpose(1, 2) # 1, t*32, 1536
pad_len = int(humo_audio_seq_len - humo_audio_input.size(1))
if pad_len > 0:
humo_audio_input = torch.nn.functional.pad(humo_audio_input, (0, 0, 0, pad_len))
should_calc = True
#TeaCache
if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
@@ -2392,25 +2458,21 @@ class WanModel(torch.nn.Module):
original_seq_len=self.original_seq_len,
enhance_enabled=enhance_enabled,
audio_scale=audio_scale,
block_mask=self.block_mask,
nag_params=nag_params,
nag_context=nag_context,
nag_params=nag_params, nag_context=nag_context,
is_uncond = is_uncond,
multitalk_audio_embedding=multitalk_audio_embedding if multitalk_audio is not None else None,
ref_target_masks=token_ref_target_masks if multitalk_audio is not None else None,
human_num=human_num if multitalk_audio is not None else 0,
inner_t=inner_t,
inner_c=inner_c,
inner_t=inner_t, inner_c=inner_c,
cross_freqs=self.cross_freqs if inner_t is not None else None,
freqs_ip=freqs_ip if x_ip is not None else None,
e_ip=e0_ip if x_ip is not None else None,
adapter_proj=adapter_proj,
ip_scale=ip_scale,
reverse_time=reverse_time,
mtv_motion_tokens=mtv_motion_tokens,
mtv_motion_rotary_emb=mtv_motion_rotary_emb,
mtv_strength=mtv_strength,
mtv_freqs=mtv_freqs
mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength, mtv_freqs=mtv_freqs,
humo_audio_input=humo_audio_input,
humo_audio_scale=humo_audio_scale,
)
if vace_data is not None: