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:
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user