From 58104b620f6b52de2fc516ade9f31383b0e7eca9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 18 Jun 2025 15:41:53 +0300 Subject: [PATCH] init --- __init__.py | 4 + fantasytalking/nodes.py | 27 ++- multitalk/multitalk.py | 374 ++++++++++++++++++++++++++++++++++++++ multitalk/nodes.py | 149 +++++++++++++++ multitalk/wav2vec2.py | 141 ++++++++++++++ nodes.py | 131 ++++++++++--- wanvideo/modules/model.py | 79 ++++++-- 7 files changed, 860 insertions(+), 45 deletions(-) create mode 100644 multitalk/multitalk.py create mode 100644 multitalk/nodes.py create mode 100644 multitalk/wav2vec2.py diff --git a/__init__.py b/__init__.py index 8035acd..31c40c5 100644 --- a/__init__.py +++ b/__init__.py @@ -7,6 +7,7 @@ from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPI from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS +from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS #from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS @@ -18,6 +19,8 @@ NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(MULTITALK_NODE_CLASS_MAPPINGS) + #NODE_CLASS_MAPPINGS.update(CAUSVID_NODE_CLASS_MAPPINGS) @@ -29,6 +32,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(ATI_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(MULTITALK_NODE_DISPLAY_NAME_MAPPINGS) #NODE_DISPLAY_NAME_MAPPINGS.update(CAUSVID_NODE_DISPLAY_NAME_MAPPINGS) diff --git a/fantasytalking/nodes.py b/fantasytalking/nodes.py index b353741..a252f5e 100644 --- a/fantasytalking/nodes.py +++ b/fantasytalking/nodes.py @@ -18,12 +18,18 @@ class DownloadAndLoadWav2VecModel: def INPUT_TYPES(s): return { "required": { - "model": (["facebook/wav2vec2-base-960h"],), + "model": ( + [ + "facebook/wav2vec2-base-960h", + "TencentGameMate/chinese-wav2vec2-base" + ], + ), "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 = ("WAV2VECMODEL",) RETURN_NAMES = ("wav2vec_model", ) @@ -31,7 +37,8 @@ class DownloadAndLoadWav2VecModel: CATEGORY = "WanVideoWrapper" def loadmodel(self, model, base_precision, load_device): - from transformers import Wav2Vec2Model, Wav2Vec2Processor + from transformers import Wav2Vec2Model, Wav2Vec2Processor, Wav2Vec2FeatureExtractor + from ..multitalk.wav2vec2 import Wav2Vec2Model as MultiTalkWav2Vec2Model 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] device = mm.get_torch_device() @@ -46,18 +53,26 @@ class DownloadAndLoadWav2VecModel: if not os.path.exists(model_path): log.info(f"Downloading Qwen model to: {model_path}") from huggingface_hub import snapshot_download + ignore_patterns = None + if model == "facebook/wav2vec2-base-960h": + ignore_patterns = ["*.bin", "*.h5"] snapshot_download( repo_id=model, - ignore_patterns=["*.bin", "*.h5"], + ignore_patterns=ignore_patterns, local_dir=model_path, local_dir_use_symlinks=False, ) - wav2vec_processor = Wav2Vec2Processor.from_pretrained(model_path) - wav2vec = Wav2Vec2Model.from_pretrained(model_path).to(base_dtype).to(transfomer_load_device).eval() + if model == "facebook/wav2vec2-base-960h": + wav2vec_processor = Wav2Vec2Processor.from_pretrained(model_path) + wav2vec = Wav2Vec2Model.from_pretrained(model_path).to(base_dtype).to(transfomer_load_device).eval() + elif model == "TencentGameMate/chinese-wav2vec2-base": + wav2vec = MultiTalkWav2Vec2Model.from_pretrained(model_path).to(base_dtype).to(transfomer_load_device).eval() + wav2vec_feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(model_path, local_files_only=True) wav2vec_processor_model = { - "processor": wav2vec_processor, + "processor": wav2vec_processor if model == "facebook/wav2vec2-base-960h" else None, + "feature_extractor": wav2vec_feature_extractor if model == "TencentGameMate/chinese-wav2vec2-base" else None, "model": wav2vec, "dtype": base_dtype,} diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py new file mode 100644 index 0000000..297cd7a --- /dev/null +++ b/multitalk/multitalk.py @@ -0,0 +1,374 @@ +from diffusers import ModelMixin, ConfigMixin +from einops import rearrange, repeat +import torch +import torch.nn as nn +from functools import lru_cache + +from comfy import model_management as mm + +def normalize_and_scale(column, source_range, target_range, epsilon=1e-8): + + source_min, source_max = source_range + new_min, new_max = target_range + + normalized = (column - source_min) / (source_max - source_min + epsilon) + scaled = normalized * (new_max - new_min) + new_min + return scaled + +def rotate_half(x): + x = rearrange(x, "... (d r) -> ... d r", r=2) + x1, x2 = x.unbind(dim=-1) + x = torch.stack((-x2, x1), dim=-1) + return rearrange(x, "... d r -> ... (d r)") + +def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None): + + ref_k = ref_k.to(visual_q.dtype).to(visual_q.device) + scale = 1.0 / visual_q.shape[-1] ** 0.5 + visual_q = visual_q * scale + visual_q = visual_q.transpose(1, 2) + ref_k = ref_k.transpose(1, 2) + attn = visual_q @ ref_k.transpose(-2, -1) + + if attn_bias is not None: + attn = attn + attn_bias + + x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens + + + x_ref_attn_maps = [] + ref_target_masks = ref_target_masks.to(visual_q.dtype) + x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype) + + for class_idx, ref_target_mask in enumerate(ref_target_masks): + mm.soft_empty_cache() + ref_target_mask = ref_target_mask[None, None, None, ...] + x_ref_attnmap = x_ref_attn_map_source * ref_target_mask + x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens + x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H + + if mode == 'mean': + x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens + elif mode == 'max': + x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens + + x_ref_attn_maps.append(x_ref_attnmap) + + del attn + del x_ref_attn_map_source + mm.soft_empty_cache() + + return torch.concat(x_ref_attn_maps, dim=0) + +def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2, enable_sp=False): + """Args: + query (torch.tensor): B M H K + key (torch.tensor): B M H K + shape (tuple): (N_t, N_h, N_w) + ref_target_masks: [B, N_h * N_w] + """ + + N_t, N_h, N_w = shape + + x_seqlens = N_h * N_w + ref_k = ref_k[:, :x_seqlens] + _, seq_lens, heads, _ = visual_q.shape + class_num, _ = ref_target_masks.shape + x_ref_attn_maps = torch.zeros(class_num, seq_lens).to(visual_q.device).to(visual_q.dtype) + + split_chunk = heads // split_num + + for i in range(split_num): + x_ref_attn_maps_perhead = calculate_x_ref_attn_map(visual_q[:, :, i*split_chunk:(i+1)*split_chunk, :], ref_k[:, :, i*split_chunk:(i+1)*split_chunk, :], ref_target_masks) + x_ref_attn_maps += x_ref_attn_maps_perhead + + return x_ref_attn_maps / split_num + +class RotaryPositionalEmbedding1D(nn.Module): + + def __init__(self, + head_dim, + ): + super().__init__() + self.head_dim = head_dim + self.base = 10000 + + + @lru_cache(maxsize=32) + def precompute_freqs_cis_1d(self, pos_indices): + + freqs = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2)[: (self.head_dim // 2)].float() / self.head_dim)) + freqs = freqs.to(pos_indices.device) + freqs = torch.einsum("..., f -> ... f", pos_indices.float(), freqs) + freqs = repeat(freqs, "... n -> ... (n r)", r=2) + return freqs + + def forward(self, x, pos_indices): + """1D RoPE. + + Args: + query (torch.tensor): [B, head, seq, head_dim] + pos_indices (torch.tensor): [seq,] + Returns: + query with the same shape as input. + """ + freqs_cis = self.precompute_freqs_cis_1d(pos_indices) + + x_ = x.float() + + freqs_cis = freqs_cis.float().to(x.device) + cos, sin = freqs_cis.cos(), freqs_cis.sin() + cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d') + x_ = (x_ * cos) + (rotate_half(x_) * sin) + + return x_.type_as(x) + +class AudioProjModel(ModelMixin, ConfigMixin): + def __init__( + self, + seq_len=5, + seq_len_vf=12, + blocks=12, + channels=768, + intermediate_dim=512, + output_dim=768, + context_tokens=32, + norm_output_audio=False, + ): + super().__init__() + + self.seq_len = seq_len + self.blocks = blocks + self.channels = channels + self.input_dim = seq_len * blocks * channels + self.input_dim_vf = seq_len_vf * blocks * channels + self.intermediate_dim = intermediate_dim + self.context_tokens = context_tokens + self.output_dim = output_dim + + # define multiple linear layers + self.proj1 = nn.Linear(self.input_dim, intermediate_dim) + self.proj1_vf = nn.Linear(self.input_dim_vf, intermediate_dim) + self.proj2 = nn.Linear(intermediate_dim, intermediate_dim) + self.proj3 = nn.Linear(intermediate_dim, context_tokens * output_dim) + self.norm = nn.LayerNorm(output_dim) if norm_output_audio else nn.Identity() + + def forward(self, audio_embeds, audio_embeds_vf): + video_length = audio_embeds.shape[1] + audio_embeds_vf.shape[1] + B, _, _, S, C = audio_embeds.shape + + # process audio of first frame + 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) + + # process audio of latter frame + audio_embeds_vf = rearrange(audio_embeds_vf, "bz f w b c -> (bz f) w b c") + batch_size_vf, window_size_vf, blocks_vf, channels_vf = audio_embeds_vf.shape + audio_embeds_vf = audio_embeds_vf.view(batch_size_vf, window_size_vf * blocks_vf * channels_vf) + + # first projection + audio_embeds = torch.relu(self.proj1(audio_embeds)) + audio_embeds_vf = torch.relu(self.proj1_vf(audio_embeds_vf)) + audio_embeds = rearrange(audio_embeds, "(bz f) c -> bz f c", bz=B) + audio_embeds_vf = rearrange(audio_embeds_vf, "(bz f) c -> bz f c", bz=B) + audio_embeds_c = torch.concat([audio_embeds, audio_embeds_vf], dim=1) + batch_size_c, N_t, C_a = audio_embeds_c.shape + audio_embeds_c = audio_embeds_c.view(batch_size_c*N_t, C_a) + + # second projection + audio_embeds_c = torch.relu(self.proj2(audio_embeds_c)) + + context_tokens = self.proj3(audio_embeds_c).reshape(batch_size_c*N_t, self.context_tokens, self.output_dim) + + # normalization and reshape + context_tokens = self.norm(context_tokens) + context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length) + + return context_tokens + +class SingleStreamAttention(nn.Module): + def __init__( + self, + dim: int, + encoder_hidden_states_dim: int, + num_heads: int, + qkv_bias: bool, + qk_norm: bool, + norm_layer: nn.Module, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + eps: float = 1e-6, + ) -> None: + super().__init__() + assert dim % num_heads == 0, "dim should be divisible by num_heads" + self.dim = dim + self.encoder_hidden_states_dim = encoder_hidden_states_dim + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.scale = self.head_dim**-0.5 + self.qk_norm = qk_norm + + self.q_linear = nn.Linear(dim, dim, bias=qkv_bias) + + self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim,eps=eps) if qk_norm else nn.Identity() + + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(dim, dim) + self.proj_drop = nn.Dropout(proj_drop) + + self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias) + + self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() + self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() + + def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: + + N_t, N_h, N_w = shape + if not enable_sp: + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + + # get q for hidden_state + B, N, C = x.shape + q = self.q_linear(x) + q_shape = (B, N, self.num_heads, self.head_dim) + q = q.view(q_shape).permute((0, 2, 1, 3)) + + if self.qk_norm: + q = self.q_norm(q) + + # get kv from encoder_hidden_states + _, N_a, _ = encoder_hidden_states.shape + encoder_kv = self.kv_linear(encoder_hidden_states) + encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim) + encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4)) + encoder_k, encoder_v = encoder_kv.unbind(0) + + if self.qk_norm: + encoder_k = self.add_k_norm(encoder_k) + + x = torch.nn.functional.scaled_dot_product_attention( + q, encoder_k, encoder_v, attn_mask=None, is_causal=False, dropout_p=0.0) + + # linear transform + x_output_shape = (B, N, C) + x = x.transpose(1, 2) + x = x.reshape(x_output_shape) + x = self.proj(x) + x = self.proj_drop(x) + + if not enable_sp: + # reshape x to origin shape + x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + + return x + +class SingleStreamMultiAttention(SingleStreamAttention): + def __init__( + self, + dim: int, + encoder_hidden_states_dim: int, + num_heads: int, + qkv_bias: bool, + qk_norm: bool, + norm_layer: nn.Module, + attn_drop: float = 0.0, + proj_drop: float = 0.0, + eps: float = 1e-6, + class_range: int = 24, + class_interval: int = 4, + ) -> None: + super().__init__( + dim=dim, + encoder_hidden_states_dim=encoder_hidden_states_dim, + num_heads=num_heads, + qkv_bias=qkv_bias, + qk_norm=qk_norm, + norm_layer=norm_layer, + attn_drop=attn_drop, + proj_drop=proj_drop, + eps=eps, + ) + self.class_interval = class_interval + self.class_range = class_range + self.rope_h1 = (0, self.class_interval) + self.rope_h2 = (self.class_range - self.class_interval, self.class_range) + self.rope_bak = int(self.class_range // 2) + + self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim) + + def forward(self, + x: torch.Tensor, + encoder_hidden_states: torch.Tensor, + shape=None, + x_ref_attn_map=None, + human_num=None) -> torch.Tensor: + + encoder_hidden_states = encoder_hidden_states.squeeze(0) + if human_num == 1: + return super().forward(x, encoder_hidden_states, shape) + + N_t, _, _ = shape + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + + # get q for hidden_state + B, N, C = x.shape + q = self.q_linear(x) + q_shape = (B, N, self.num_heads, self.head_dim) + q = q.view(q_shape).permute((0, 2, 1, 3)) + + if self.qk_norm: + q = self.q_norm(q) + + + max_values = x_ref_attn_map.max(1).values[:, None, None] + min_values = x_ref_attn_map.min(1).values[:, None, None] + max_min_values = torch.cat([max_values, min_values], dim=2) + + human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min() + human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min() + + human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1])) + human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1])) + back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device) + max_indices = x_ref_attn_map.argmax(dim=0) + normalized_map = torch.stack([human1, human2, back], dim=1) + normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices] # N + + q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t) + q = self.rope_1d(q, normalized_pos) + q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t) + + _, N_a, _ = encoder_hidden_states.shape + encoder_kv = self.kv_linear(encoder_hidden_states) + encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim) + encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4)) + encoder_k, encoder_v = encoder_kv.unbind(0) + + if self.qk_norm: + encoder_k = self.add_k_norm(encoder_k) + + + per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device) + per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2 + per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2 + encoder_pos = torch.concat([per_frame]*N_t, dim=0) + encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t) + encoder_k = self.rope_1d(encoder_k, encoder_pos) + encoder_k = rearrange(encoder_k, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t) + + x = torch.nn.functional.scaled_dot_product_attention( + q, encoder_k, encoder_v, attn_mask=None, is_causal=False, dropout_p=0.0) + + # linear transform + x_output_shape = (B, N, C) + x = x.transpose(1, 2) + x = x.reshape(x_output_shape) + x = self.proj(x) + x = self.proj_drop(x) + + # reshape x to origin shape + x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + + return x \ No newline at end of file diff --git a/multitalk/nodes.py b/multitalk/nodes.py new file mode 100644 index 0000000..f1d8d68 --- /dev/null +++ b/multitalk/nodes.py @@ -0,0 +1,149 @@ +import folder_paths +from comfy import model_management as mm +from comfy.utils import load_torch_file +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device +import torch + +class MultiTalkModelLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), + + "base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}), + }, + } + + RETURN_TYPES = ("MULTITALKMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, base_precision): + from .multitalk import AudioProjModel + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + 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] + + model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) + sd = load_torch_file(model_path, device=offload_device, safe_load=True) + + audio_proj_keys = [k for k in sd.keys() if "audio_proj" in k] + audio_proj_sd = {k.replace("audio_proj.", ""): sd.pop(k) for k in audio_proj_keys} + + audio_window=5 + intermediate_dim=512 + output_dim=768 + context_tokens=32 + vae_scale=4 + norm_output_audio = True + + with init_empty_weights(): + multitalk_proj_model = AudioProjModel( + seq_len=audio_window, + seq_len_vf=audio_window+vae_scale-1, + intermediate_dim=intermediate_dim, + output_dim=output_dim, + context_tokens=context_tokens, + norm_output_audio=norm_output_audio, + ) + #fantasytalking_proj_model.load_state_dict(sd, strict=False) + + for name, param in multitalk_proj_model.named_parameters(): + set_module_tensor_to_device(multitalk_proj_model, name, device=offload_device, dtype=base_dtype, value=audio_proj_sd[name]) + + multitalk = { + "proj_model": multitalk_proj_model, + "sd": sd, + } + + return (multitalk,) + +class MultiTalkWav2VecEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "wav2vec_model": ("WAV2VECMODEL",), + "audio": ("AUDIO",), + "num_frames": ("INT", {"default": 81, "min": 1, "max": 1000, "step": 1}), + "fps": ("FLOAT", {"default": 23.0, "min": 1.0, "max": 60.0, "step": 0.1}), + "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}), + "audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}), + }, + } + + RETURN_TYPES = ("MULTITALK_EMBEDS", ) + RETURN_NAMES = ("multitalk_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, wav2vec_model, fps, num_frames, audio, audio_scale, audio_cfg_scale): + import torchaudio + import numpy as np + from einops import rearrange + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + dtype = wav2vec_model["dtype"] + wav2vec = wav2vec_model["model"] + wav2vec_feature_extractor = wav2vec_model["feature_extractor"] + + sr = 16000 + + audio_input = audio["waveform"] + sample_rate = audio["sample_rate"] + if sample_rate != sr: + audio_input = torchaudio.functional.resample(audio_input, sample_rate, sr) + audio_input = audio_input[0][0] + + start_time = 0 + end_time = num_frames / fps + + start_sample = int(start_time * sr) + end_sample = int(end_time * sr) + + try: + audio_segment = audio_input[start_sample:end_sample] + except: + audio_segment = audio_input + + print("audio_segment.shape", audio_segment.shape) + + audio_feature = np.squeeze( + wav2vec_feature_extractor(audio_segment.numpy(), sampling_rate=sr).input_values + ) + + audio_feature = torch.from_numpy(audio_feature).float().to(device=device) + audio_feature = audio_feature.unsqueeze(0) + + # audio encoder + embeddings = wav2vec(audio_feature.to(dtype), seq_len=int(num_frames), output_hidden_states=True) + + if len(embeddings) == 0: + print("Fail to extract audio embedding") + return None + + audio_emb = torch.stack(embeddings.hidden_states[1:], dim=1).squeeze(0) + audio_emb = rearrange(audio_emb, "b s d -> s b d") + + multitalk_embeds = { + "audio_features": audio_emb, + "audio_scale": audio_scale, + "audio_cfg_scale": audio_cfg_scale + } + + return (multitalk_embeds,) + +NODE_CLASS_MAPPINGS = { + "MultiTalkModelLoader": MultiTalkModelLoader, + "MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds, + +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MultiTalkModelLoader": "MultiTalk Model Loader", + "MultiTalkWav2VecEmbeds": "MultiTalk Wav2Vec Embeds", +} \ No newline at end of file diff --git a/multitalk/wav2vec2.py b/multitalk/wav2vec2.py new file mode 100644 index 0000000..c0b0af2 --- /dev/null +++ b/multitalk/wav2vec2.py @@ -0,0 +1,141 @@ +from transformers import Wav2Vec2Config, Wav2Vec2Model +from transformers.modeling_outputs import BaseModelOutput +import torch +import torch.nn.functional as F + +def get_mask_from_lengths(lengths, max_len=None): + lengths = lengths.to(torch.long) + if max_len is None: + max_len = torch.max(lengths).item() + + ids = torch.arange(0, max_len).unsqueeze(0).expand(lengths.shape[0], -1).to(lengths.device) + mask = ids < lengths.unsqueeze(1).expand(-1, max_len) + + return mask + + +def linear_interpolation(features, seq_len): + features = features.transpose(1, 2) + output_features = F.interpolate(features, size=seq_len, align_corners=True, mode='linear') + return output_features.transpose(1, 2) + +# the implementation of Wav2Vec2Model is borrowed from +# https://github.com/huggingface/transformers/blob/HEAD/src/transformers/models/wav2vec2/modeling_wav2vec2.py +# initialize our encoder with the pre-trained wav2vec 2.0 weights. +class Wav2Vec2Model(Wav2Vec2Model): + def __init__(self, config: Wav2Vec2Config): + super().__init__(config) + + def forward( + self, + input_values, + seq_len, + attention_mask=None, + mask_time_indices=None, + output_attentions=None, + output_hidden_states=None, + return_dict=None, + ): + self.config.output_attentions = True + + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + extract_features = self.feature_extractor(input_values) + extract_features = extract_features.transpose(1, 2) + extract_features = linear_interpolation(extract_features, seq_len=seq_len) + + if attention_mask is not None: + # compute reduced attention_mask corresponding to feature vectors + attention_mask = self._get_feature_vector_attention_mask( + extract_features.shape[1], attention_mask, add_adapter=False + ) + + hidden_states, extract_features = self.feature_projection(extract_features) + hidden_states = self._mask_hidden_states( + hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask + ) + + encoder_outputs = self.encoder( + hidden_states, + attention_mask=attention_mask, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + hidden_states = encoder_outputs[0] + + if self.adapter is not None: + hidden_states = self.adapter(hidden_states) + + if not return_dict: + return (hidden_states, ) + encoder_outputs[1:] + return BaseModelOutput( + last_hidden_state=hidden_states, + hidden_states=encoder_outputs.hidden_states, + attentions=encoder_outputs.attentions, + ) + + + def feature_extract( + self, + input_values, + seq_len, + ): + extract_features = self.feature_extractor(input_values) + extract_features = extract_features.transpose(1, 2) + extract_features = linear_interpolation(extract_features, seq_len=seq_len) + + return extract_features + + def encode( + self, + extract_features, + attention_mask=None, + mask_time_indices=None, + output_attentions=None, + output_hidden_states=None, + return_dict=None, + ): + self.config.output_attentions = True + + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + if attention_mask is not None: + # compute reduced attention_mask corresponding to feature vectors + attention_mask = self._get_feature_vector_attention_mask( + extract_features.shape[1], attention_mask, add_adapter=False + ) + + + hidden_states, extract_features = self.feature_projection(extract_features) + hidden_states = self._mask_hidden_states( + hidden_states, mask_time_indices=mask_time_indices, attention_mask=attention_mask + ) + + encoder_outputs = self.encoder( + hidden_states, + attention_mask=attention_mask, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + hidden_states = encoder_outputs[0] + + if self.adapter is not None: + hidden_states = self.adapter(hidden_states) + + if not return_dict: + return (hidden_states, ) + encoder_outputs[1:] + return BaseModelOutput( + last_hidden_state=hidden_states, + hidden_states=encoder_outputs.hidden_states, + attentions=encoder_outputs.attentions, + ) diff --git a/nodes.py b/nodes.py index b20232b..1ab9f4e 100644 --- a/nodes.py +++ b/nodes.py @@ -538,6 +538,7 @@ class WanVideoModelLoader: "vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}), "vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}), "fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}), + "multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}), } } @@ -547,7 +548,7 @@ class WanVideoModelLoader: CATEGORY = "WanVideoWrapper" def loadmodel(self, model, base_precision, load_device, quantization, - compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None): + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None): assert not (vram_management_args is not None and block_swap_args is not None), "Can't use both block_swap_args and vram_management_args at the same time" lora_low_mem_load = False if lora is not None: @@ -733,8 +734,31 @@ class WanVideoModelLoader: block.cross_attn.k_proj = nn.Linear(context_dim, dim, bias=False) block.cross_attn.v_proj = nn.Linear(context_dim, dim, bias=False) sd.update(fantasytalking_model["sd"]) + if multitalk_model is not None: + # init audio module + from .multitalk.multitalk import SingleStreamMultiAttention + from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm + norm_input_visual = True #dunno what this is + + for block in transformer.blocks: + block.audio_cross_attn = SingleStreamMultiAttention( + dim=dim, + encoder_hidden_states_dim=768, + num_heads=num_heads, + qk_norm=False, + qkv_bias=True, + eps=transformer.eps, + norm_layer=WanRMSNorm, + class_range=24, + class_interval=4 + ) + block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity() + log.info("MultiTalk model detected, patching model...") + + sd.update(multitalk_model["sd"]) + - # RealisDance-DiT + # Additional cond latents if "add_conv_in.weight" in sd: def zero_module(module): for p in module.parameters(): @@ -856,6 +880,9 @@ class WanVideoModelLoader: del sd + if multitalk_model is not None: + transformer.audio_proj = multitalk_model["proj_model"] + if vram_management_args is not None: from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm @@ -1714,8 +1741,8 @@ class WanVideoRealisDanceLatents: }, } - RETURN_TYPES = ("REALISDANCELATENTS",) - RETURN_NAMES = ("realisdance_latents",) + RETURN_TYPES = ("ADD_COND_LATENTS",) + RETURN_NAMES = ("add_cond_latents",) FUNCTION = "process" CATEGORY = "WanVideoWrapper" @@ -1727,14 +1754,14 @@ class WanVideoRealisDanceLatents: pose_latent = torch.cat((smpl_latent["samples"], hamer), dim=1) - realisdance_latents = { + add_cond_latents = { "ref_latent": ref_latent["samples"], "pose_latent": pose_latent, "pose_cond_start_percent": pose_cond_start_percent, "pose_cond_end_percent": pose_cond_end_percent, } - return (realisdance_latents,) + return (add_cond_latents,) class WanVideoImageToVideoEncode: @classmethod @@ -1758,7 +1785,7 @@ class WanVideoImageToVideoEncode: "temporal_mask": ("MASK", {"tooltip": "mask"}), "extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), - "realisdance_latents": ("REALISDANCELATENTS", {"tooltip": "RealisDance latents"}), + "add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}), } } @@ -1769,7 +1796,7 @@ class WanVideoImageToVideoEncode: def process(self, vae, width, height, num_frames, force_offload, noise_aug_strength, start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, - temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, realisdance_latents=None): + temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1872,8 +1899,8 @@ class WanVideoImageToVideoEncode: frames_per_stride = (num_frames - 1) // 4 + (2 if end_image is not None and not fun_or_fl2v_model else 1) max_seq_len = frames_per_stride * patches_per_frame - if realisdance_latents is not None: - realisdance_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device) + if add_cond_latents is not None: + add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device) vae.model.clear_cache() if force_offload: @@ -1893,7 +1920,7 @@ class WanVideoImageToVideoEncode: "end_image": resized_end_image if end_image is not None else None, "fun_or_fl2v_model": fun_or_fl2v_model, "has_ref": has_ref, - "realisdance_latents": realisdance_latents + "add_cond_latents": add_cond_latents } return (image_embeds,) @@ -2531,6 +2558,7 @@ class WanVideoSampler: "unianimate_poses": ("UNIANIMATE_POSE", ), "fantasytalking_embeds": ("FANTASYTALKING_EMBEDS", ), "uni3c_embeds": ("UNI3C_EMBEDS", ), + "multitalk_embeds": ("MULTITALK_EMBEDS", ), } } @@ -2542,7 +2570,7 @@ class WanVideoSampler: def process(self, model, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, text_embeds=None, force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None, cache_args=None, teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, - experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None): + experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None, multitalk_embeds=None): patcher = model model = model.model @@ -2676,13 +2704,13 @@ class WanVideoSampler: image_cond_ati = patch_motion(ATI_tracks.to(image_cond.device, image_cond.dtype), image_cond, topk=topk, temperature=temperature) log.info(f"ATI tracks shape: {ATI_tracks.shape}") - realisdance_latents = image_embeds.get("realisdance_latents", None) - if realisdance_latents is not None: - add_cond = realisdance_latents["pose_latent"] - attn_cond = realisdance_latents["ref_latent"] - attn_cond_neg = realisdance_latents["ref_latent_neg"] - add_cond_start_percent = realisdance_latents["pose_cond_start_percent"] - add_cond_end_percent = realisdance_latents["pose_cond_end_percent"] + add_cond_latents = image_embeds.get("add_cond_latents", None) + if add_cond_latents is not None: + add_cond = add_cond_latents["pose_latent"] + attn_cond = add_cond_latents["ref_latent"] + attn_cond_neg = add_cond_latents["ref_latent_neg"] + add_cond_start_percent = add_cond_latents["pose_cond_start_percent"] + add_cond_end_percent = add_cond_latents["pose_cond_end_percent"] end_image = image_embeds.get("end_image", None) lat_h = image_embeds.get("lat_h", None) @@ -2855,7 +2883,8 @@ class WanVideoSampler: "end_percent": unianimate_poses["end_percent"] } - audio_proj = None + audio_proj = multitalk_audio_embedding = None + audio_scale = 1.0 if fantasytalking_embeds is not None: audio_proj = fantasytalking_embeds["audio_proj"].to(device) audio_context_lens = fantasytalking_embeds["audio_context_lens"] @@ -2864,6 +2893,14 @@ class WanVideoSampler: if not isinstance(audio_cfg_scale, list): audio_cfg_scale = [audio_cfg_scale] * (steps +1) log.info(f"Audio proj shape: {audio_proj.shape}, audio context lens: {audio_context_lens}") + elif multitalk_embeds is not None: + multitalk_audio_embedding = multitalk_embeds["audio_features"].to(device, dtype) + audio_scale = multitalk_embeds["audio_scale"] + audio_cfg_scale = multitalk_embeds["audio_cfg_scale"] + if not isinstance(audio_cfg_scale, list): + audio_cfg_scale = [audio_cfg_scale] * (steps +1) + log.info(f"Multitalk audio features shape: {multitalk_audio_embedding.shape}") + minimax_latents = minimax_mask_latents = None minimax_latents = image_embeds.get("minimax_latents", None) @@ -3240,6 +3277,25 @@ class WanVideoSampler: if minimax_latents is not None: z_pos = z_neg = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0) + + if multitalk_audio_embedding is not None: + audio_embedding = [multitalk_audio_embedding] + audio_embs = [] + indices = (torch.arange(4 + 1) - 2) * 1 + # split audio with window size + for human_idx in range(1): + center_indices = torch.arange( + 0, + latent_video_length * 4 + 1 if add_cond is not None else (latent_video_length-1) * 4 + 1, + 1, + ).unsqueeze( + 1 + ) + indices.unsqueeze(0) + center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1) + audio_emb = audio_embedding[human_idx][center_indices][None,...].to(device) + audio_embs.append(audio_emb) + audio_embs = torch.concat(audio_embs, dim=0).to(dtype) + print("audio_embs: ", audio_embs.shape) base_params = { 'seq_len': seq_len, @@ -3254,12 +3310,13 @@ class WanVideoSampler: 'fun_camera': control_camera_input if control_camera_latents is not None else None, 'audio_proj': audio_proj if fantasytalking_embeds is not None else None, 'audio_context_lens': audio_context_lens if fantasytalking_embeds is not None else None, - 'audio_scale': audio_scale if fantasytalking_embeds is not None else None, + 'audio_scale': audio_scale, "pcd_data": pcd_data, "controlnet": controlnet, "add_cond": add_cond_input, "nag_params": text_embeds.get("nag_params", {}), "nag_context": text_embeds.get("nag_prompt_embeds", None), + "multitalk_audio": audio_embs if multitalk_audio_embedding is not None else None, } batch_size = 1 @@ -3333,6 +3390,25 @@ class WanVideoSampler: + 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'] = None + noise_pred_no_audio, cache_state_audio = transformer( + [z_pos], 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_uncond + + cfg_scale * (noise_pred_cond - noise_pred_uncond) + + audio_cfg_scale[idx] * (noise_pred_uncond - noise_pred_no_audio) + ) + return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_audio] #batched else: @@ -3776,6 +3852,9 @@ class WanVideoDecode: "tile_stride_x": ("INT", {"default": 144, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride width in pixels. Smaller values use less VRAM but will introduce more seams."}), "tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}), }, + "optional": { + "normalization": (["default", "minmax"], {"advanced": True}), + } } @classmethod @@ -3791,7 +3870,7 @@ class WanVideoDecode: FUNCTION = "decode" CATEGORY = "WanVideoWrapper" - def decode(self, vae, samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y): + def decode(self, vae, samples, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, normalization="default"): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.soft_empty_cache() @@ -3825,9 +3904,11 @@ class WanVideoDecode: images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0] vae.model.clear_cache() - #images = (images - images.min()) / (images.max() - images.min()) - images = torch.clamp(images, -1.0, 1.0) - images = (images + 1.0) / 2.0 + if normalization == "minmax": + images = (images - images.min()) / (images.max() - images.min()) + else: + images = torch.clamp(images, -1.0, 1.0) + images = (images + 1.0) / 2.0 if is_looped: #images = images[:, warmup_latent_count * 4:] diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 7040120..c65b959 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -26,6 +26,8 @@ import gc import comfy.model_management as mm from ...utils import log, get_module_memory_mb +from ...multitalk.multitalk import get_attn_map_with_target + from comfy.ldm.flux.math import apply_rope as apply_rope_comfy def rope_riflex(pos, dim, theta, L_test, k, temporal): @@ -188,7 +190,7 @@ class WanSelfAttention(nn.Module): self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity() - def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default", block_mask=None): + def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default", block_mask=None, ref_target_masks=None): r""" Args: x(Tensor): Shape [B, L, num_heads, C / num_heads] @@ -267,7 +269,12 @@ class WanSelfAttention(nn.Module): if is_enhance_enabled(): x *= feta_scores - return x + #multitalk + x_ref_attn_map = None + if ref_target_masks is not None: + x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0], + ref_target_masks=ref_target_masks) + return x, x_ref_attn_map def forward_split(self, x, seq_lens, grid_sizes, freqs, seq_chunks=1,current_step=0, video_attention_split_steps = [], rope_func = "default"): r""" @@ -586,7 +593,10 @@ class WanAttentionBlock(nn.Module): block_mask=None, nag_params={}, nag_context=None, - is_uncond=False + is_uncond=False, + multitalk_audio_embedding=None, + ref_target_masks=None, + human_num=0 ): r""" Args: @@ -620,11 +630,12 @@ class WanAttentionBlock(nn.Module): video_attention_split_steps=video_attention_split_steps ) else: - y = self.self_attn.forward( + y, x_ref_attn_map = self.self_attn.forward( input_x, seq_lens, grid_sizes, freqs, rope_func=rope_func, block_mask=block_mask, + ref_target_masks=ref_target_masks, ) #ReCamMaster if camera_embed is not None: @@ -644,19 +655,26 @@ class WanAttentionBlock(nn.Module): else: x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes, audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, - num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond) + num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, + multitalk_audio_embedding=multitalk_audio_embedding, x_ref_attn_map=x_ref_attn_map, human_num=human_num) else: y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) x = x + (y * e[5]) + del e return x #@torch.compiler.disable() def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, nag_params={}, - nag_context=None, is_uncond=False): + nag_context=None, is_uncond=False, multitalk_audio_embedding=None, x_ref_attn_map=None, human_num=0): x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed, audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond) + #multitalk + if multitalk_audio_embedding is not None: + x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding, + shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num) + x = x + x_audio * audio_scale y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) x = x + (y * e[5]) return x @@ -1223,7 +1241,9 @@ class WanModel(ModelMixin, ConfigMixin): add_cond=None, attn_cond=None, nag_params={}, - nag_context=None + nag_context=None, + multitalk_audio=None, + ref_target_masks=None ): r""" Forward pass through the diffusion model @@ -1367,9 +1387,9 @@ class WanModel(ModelMixin, ConfigMixin): # time embeddings if t.dim() == 2: b, f = t.shape - _flag_df = True + diffusion_forcing = True else: - _flag_df = False + diffusion_forcing = False e = self.time_embedding( sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(x.dtype) @@ -1380,12 +1400,12 @@ class WanModel(ModelMixin, ConfigMixin): fps_embeds = torch.tensor(fps_embeds, dtype=torch.long, device=device) fps_emb = self.fps_embedding(fps_embeds).to(e0.dtype) - if _flag_df: + if diffusion_forcing: e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)).repeat(t.shape[1], 1, 1) else: e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)) - if _flag_df: + if diffusion_forcing: e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim) e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim) @@ -1398,9 +1418,9 @@ class WanModel(ModelMixin, ConfigMixin): e = e.to(self.offload_device, non_blocking=self.use_non_blocking) - # context (test embedding) + # context (text embedding) context_lens = None - if hasattr(self, "text_embedding"): + if hasattr(self, "text_embedding") and context != []: if self.offload_txt_emb: self.text_embedding.to(self.main_device) context = self.text_embedding( @@ -1433,6 +1453,34 @@ class WanModel(ModelMixin, ConfigMixin): if self.offload_img_emb: self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) + # MultiTalk + if multitalk_audio is not None: + audio_cond = multitalk_audio.to(device=x.device, dtype=x.dtype) + first_frame_audio_emb_s = audio_cond[:, :1, ...] + latter_frame_audio_emb = audio_cond[:, 1:, ...] + latter_frame_audio_emb = rearrange(latter_frame_audio_emb, "b (n_t n) w s c -> b n_t n w s c", n=4) + middle_index = self.audio_proj.seq_len // 2 + latter_first_frame_audio_emb = latter_frame_audio_emb[:, :, :1, :middle_index+1, ...] + latter_first_frame_audio_emb = rearrange(latter_first_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c") + latter_last_frame_audio_emb = latter_frame_audio_emb[:, :, -1:, middle_index:, ...] + latter_last_frame_audio_emb = rearrange(latter_last_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c") + latter_middle_frame_audio_emb = latter_frame_audio_emb[:, :, 1:-1, middle_index:middle_index+1, ...] + latter_middle_frame_audio_emb = rearrange(latter_middle_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c") + latter_frame_audio_emb_s = torch.concat([latter_first_frame_audio_emb, latter_middle_frame_audio_emb, latter_last_frame_audio_emb], dim=2) + multitalk_audio_embedding = self.audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s) + human_num = len(multitalk_audio_embedding) + multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(x.dtype) + + # convert ref_target_masks to token_ref_target_masks + # !not implemented! + if ref_target_masks is not None: + ref_target_masks = ref_target_masks.unsqueeze(0).to(torch.float32) + token_ref_target_masks = nn.functional.interpolate(ref_target_masks, size=(H // 2, W // 2), mode='nearest') + token_ref_target_masks = token_ref_target_masks.squeeze(0) + token_ref_target_masks = (token_ref_target_masks > 0) + 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) + should_calc = True accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step: @@ -1540,7 +1588,10 @@ class WanModel(ModelMixin, ConfigMixin): block_mask=self.block_mask, nag_params=nag_params, nag_context=nag_context, - is_uncond = is_uncond + is_uncond = is_uncond, + multitalk_audio_embedding=multitalk_audio_embedding if multitalk_audio is not None else None, + ref_target_masks=ref_target_masks if multitalk_audio is not None else None, + human_num=human_num if multitalk_audio is not None else 0 ) if vace_data is not None: