From 5683d8306a75dc8ae845ecffc9f1d0e06043761b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 28 Apr 2025 19:10:32 +0300 Subject: [PATCH] Support FantasyTalking --- __init__.py | 3 + fantasytalking/model.py | 130 ++++++++++++++++++++++++++ fantasytalking/nodes.py | 191 ++++++++++++++++++++++++++++++++++++++ nodes.py | 51 +++++++++- wanvideo/modules/model.py | 47 ++++++++-- 5 files changed, 412 insertions(+), 10 deletions(-) create mode 100644 fantasytalking/model.py create mode 100644 fantasytalking/nodes.py diff --git a/__init__.py b/__init__.py index fda1074..fe65d44 100644 --- a/__init__.py +++ b/__init__.py @@ -2,13 +2,16 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS +from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_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) +NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/fantasytalking/model.py b/fantasytalking/model.py new file mode 100644 index 0000000..d7bd245 --- /dev/null +++ b/fantasytalking/model.py @@ -0,0 +1,130 @@ +import os + +import torch +import torch.nn as nn +import torch.nn.functional as F +from safetensors import safe_open + +class AudioProjModel(nn.Module): + def __init__(self, audio_in_dim=1024, cross_attention_dim=1024): + super().__init__() + self.cross_attention_dim = cross_attention_dim + self.proj = torch.nn.Linear(audio_in_dim, cross_attention_dim, bias=False) + self.norm = torch.nn.LayerNorm(cross_attention_dim) + + def forward(self, audio_embeds): + context_tokens = self.proj(audio_embeds) + context_tokens = self.norm(context_tokens) + return context_tokens # [B,L,C] + +class FantasyTalkingAudioConditionModel(nn.Module): + def __init__(self, audio_in_dim: int, audio_proj_dim: int): + super().__init__() + + self.audio_in_dim = audio_in_dim + self.audio_proj_dim = audio_proj_dim + + # audio proj model + self.proj_model = self.init_proj(self.audio_proj_dim) + + def init_proj(self, cross_attention_dim=5120): + proj_model = AudioProjModel( + audio_in_dim=self.audio_in_dim, cross_attention_dim=cross_attention_dim + ) + return proj_model + + def get_proj_fea(self, audio_fea=None): + return self.proj_model(audio_fea) if audio_fea is not None else None + + def split_audio_sequence(self, audio_proj_length, num_frames=81): + """ + Map the audio feature sequence to corresponding latent frame slices. + + Args: + audio_proj_length (int): The total length of the audio feature sequence + (e.g., 173 in audio_proj[1, 173, 768]). + num_frames (int): The number of video frames in the training data (default: 81). + + Returns: + list: A list of [start_idx, end_idx] pairs. Each pair represents the index range + (within the audio feature sequence) corresponding to a latent frame. + """ + # Average number of tokens per original video frame + tokens_per_frame = audio_proj_length / num_frames + + # Each latent frame covers 4 video frames, and we want the center + tokens_per_latent_frame = tokens_per_frame * 4 + half_tokens = int(tokens_per_latent_frame / 2) + + pos_indices = [] + for i in range(int((num_frames - 1) / 4) + 1): + if i == 0: + pos_indices.append(0) + else: + start_token = tokens_per_frame * ((i - 1) * 4 + 1) + end_token = tokens_per_frame * (i * 4 + 1) + center_token = int((start_token + end_token) / 2) - 1 + pos_indices.append(center_token) + + # Build index ranges centered around each position + pos_idx_ranges = [[idx - half_tokens, idx + half_tokens] for idx in pos_indices] + + # Adjust the first range to avoid negative start index + pos_idx_ranges[0] = [ + -(half_tokens * 2 - pos_idx_ranges[1][0]), + pos_idx_ranges[1][0], + ] + + return pos_idx_ranges + + def split_tensor_with_padding(self, input_tensor, pos_idx_ranges, expand_length=0): + """ + Split the input tensor into subsequences based on index ranges, and apply right-side zero-padding + if the range exceeds the input boundaries. + + Args: + input_tensor (Tensor): Input audio tensor of shape [1, L, 768]. + pos_idx_ranges (list): A list of index ranges, e.g. [[-7, 1], [1, 9], ..., [165, 173]]. + expand_length (int): Number of tokens to expand on both sides of each subsequence. + + Returns: + sub_sequences (Tensor): A tensor of shape [1, F, L, 768], where L is the length after padding. + Each element is a padded subsequence. + k_lens (Tensor): A tensor of shape [F], representing the actual (unpadded) length of each subsequence. + Useful for ignoring padding tokens in attention masks. + """ + pos_idx_ranges = [ + [idx[0] - expand_length, idx[1] + expand_length] for idx in pos_idx_ranges + ] + sub_sequences = [] + seq_len = input_tensor.size(1) # 173 + max_valid_idx = seq_len - 1 # 172 + k_lens_list = [] + for start, end in pos_idx_ranges: + # Calculate the fill amount + pad_front = max(-start, 0) + pad_back = max(end - max_valid_idx, 0) + + # Calculate the start and end indices of the valid part + valid_start = max(start, 0) + valid_end = min(end, max_valid_idx) + + # Extract the valid part + if valid_start <= valid_end: + valid_part = input_tensor[:, valid_start : valid_end + 1, :] + else: + valid_part = input_tensor.new_zeros((1, 0, input_tensor.size(2))) + + # In the sequence dimension (the 1st dimension) perform padding + padded_subseq = F.pad( + valid_part, + (0, 0, 0, pad_back + pad_front, 0, 0), + mode="constant", + value=0, + ) + k_lens_list.append(padded_subseq.size(-2) - pad_back - pad_front) + + sub_sequences.append(padded_subseq) + return torch.stack(sub_sequences, dim=1), torch.tensor( + k_lens_list, dtype=torch.long + ) diff --git a/fantasytalking/nodes.py b/fantasytalking/nodes.py new file mode 100644 index 0000000..2e6a865 --- /dev/null +++ b/fantasytalking/nodes.py @@ -0,0 +1,191 @@ +import os +import torch +import gc +from ..utils import log + +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device + +import comfy.model_management as mm +from comfy.utils import load_torch_file +import folder_paths + +script_directory = os.path.dirname(os.path.abspath(__file__)) + + +class DownloadAndLoadWav2VecModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (["facebook/wav2vec2-base-960h"],), + + "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", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, base_precision, load_device): + from transformers import Wav2Vec2Model, Wav2Vec2Processor + + 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() + offload_device = mm.unet_offload_device() + + if load_device == "offload_device": + transfomer_load_device = offload_device + else: + transfomer_load_device = device + + model_path = os.path.join(folder_paths.models_dir, "transformers", model) + if not os.path.exists(model_path): + log.info(f"Downloading Qwen model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download( + repo_id=model, + 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() + + wav2vec_processor_model = { + "processor": wav2vec_processor, + "model": wav2vec, + "dtype": base_dtype,} + + return (wav2vec_processor_model,) + +class FantasyTalkingModelLoader: + @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 = ("FANTASYTALKINGMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, base_precision): + from .model import FantasyTalkingAudioConditionModel + + 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) + + with init_empty_weights(): + fantasytalking_proj_model = FantasyTalkingAudioConditionModel(audio_in_dim=768, audio_proj_dim=2048) + #fantasytalking_proj_model.load_state_dict(sd, strict=False) + + for name, param in fantasytalking_proj_model.named_parameters(): + set_module_tensor_to_device(fantasytalking_proj_model, name, device=offload_device, dtype=base_dtype, value=sd[name]) + + fantasytalking = { + "proj_model": fantasytalking_proj_model, + "sd": sd, + } + + return (fantasytalking,) + +class FantasyTalkingWav2VecEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "wav2vec_model": ("WAV2VECMODEL",), + "fantasytalking_model": ("FANTASYTALKINGMODEL",), + "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 = ("FANTASYTALKING_EMBEDS", ) + RETURN_NAMES = ("fantasytalking_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, wav2vec_model, fantasytalking_model, fps, num_frames, audio_scale, audio_cfg_scale, audio): + import torchaudio + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + dtype = wav2vec_model["dtype"] + wav2vec = wav2vec_model["model"] + wav2vec_processor = wav2vec_model["processor"] + audio_proj_model = fantasytalking_model["proj_model"] + + 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) + + input_values = wav2vec_processor( + audio_segment.numpy(), sampling_rate=sr, return_tensors="pt" + ).input_values.to(dtype).to(device) + + audio_features = wav2vec(input_values).last_hidden_state + + audio_proj_model.proj_model.to(device) + audio_proj_fea = audio_proj_model.get_proj_fea(audio_features) + pos_idx_ranges = audio_proj_model.split_audio_sequence( + audio_proj_fea.size(1), num_frames=num_frames + ) + audio_proj_split, audio_context_lens = audio_proj_model.split_tensor_with_padding( + audio_proj_fea, pos_idx_ranges, expand_length=4 + ) # [b,21,9+8,768] + audio_proj_model.proj_model.to(offload_device) + mm.soft_empty_cache() + + out = { + "audio_proj": audio_proj_split, + "audio_context_lens": audio_context_lens, + "audio_scale": audio_scale, + "audio_cfg_scale": audio_cfg_scale + } + + return (out,) + + +NODE_CLASS_MAPPINGS = { + "DownloadAndLoadWav2VecModel": DownloadAndLoadWav2VecModel, + "FantasyTalkingModelLoader": FantasyTalkingModelLoader, + "FantasyTalkingWav2VecEmbeds": FantasyTalkingWav2VecEmbeds, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "DownloadAndLoadWav2VecModel": "(Down)load Wav2Vec Model", + "FantasyTalkingModelLoader": "FantasyTalking Model Loader", + "FantasyTalkingWav2VecEmbeds": "FantasyTalking Wav2Vec Embeds", + } diff --git a/nodes.py b/nodes.py index 9a0a5c9..83a5d44 100644 --- a/nodes.py +++ b/nodes.py @@ -482,6 +482,7 @@ class WanVideoModelLoader: "lora": ("WANVIDLORA", {"default": None}), "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"}), } } @@ -491,7 +492,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): + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_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: @@ -631,6 +632,7 @@ class WanVideoModelLoader: transformer = WanModel(**TRANSFORMER_CONFIG) transformer.eval() + #ReCamMaster if "blocks.0.cam_encoder.weight" in sd: log.info("ReCamMaster model detected, patching model...") import torch.nn as nn @@ -641,6 +643,16 @@ class WanVideoModelLoader: block.cam_encoder.bias.data.zero_() block.projector.weight = nn.Parameter(torch.eye(dim)) block.projector.bias = nn.Parameter(torch.zeros(dim)) + + # FantasyTalking https://github.com/Fantasy-AMAP + if fantasytalking_model is not None: + log.info("FantasyTalking model detected, patching model...") + context_dim = fantasytalking_model["sd"]["proj_model.proj.weight"].shape[0] + import torch.nn as nn + for block in transformer.blocks: + 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"]) comfy_model = WanVideoModel( WanVideoModelConfig(base_dtype), @@ -2271,6 +2283,7 @@ class WanVideoSampler: "experimental_args": ("EXPERIMENTALARGS", ), "sigmas": ("SIGMAS", ), "unianimate_poses": ("UNIANIMATE_POSE", ), + "fantasytalking_embeds": ("FANTASYTALKING_EMBEDS", ), } } @@ -2281,7 +2294,8 @@ class WanVideoSampler: def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=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): + 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): #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." patcher = model model = model.model @@ -2503,8 +2517,13 @@ class WanVideoSampler: "start_percent": unianimate_poses["start_percent"], "end_percent": unianimate_poses["end_percent"] } - - + + if fantasytalking_embeds is not None: + audio_proj = fantasytalking_embeds["audio_proj"].to(device) + audio_context_lens = fantasytalking_embeds["audio_context_lens"] + audio_scale = fantasytalking_embeds["audio_scale"] + audio_cfg_scale = fantasytalking_embeds["audio_cfg_scale"] + log.info(f"Audio proj shape: {audio_proj.shape}, audio context lens: {audio_context_lens}") is_looped = False if context_options is not None: @@ -2808,6 +2827,9 @@ class WanVideoSampler: 'camera_embed': camera_embed, 'unianim_data': unianim_data, 'fun_ref': fun_ref_input if fun_ref_image 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, } batch_size = 1 @@ -2835,6 +2857,9 @@ class WanVideoSampler: ) return noise_pred_cond, [teacache_state_cond] #uncond + if fantasytalking_embeds is not None: + if not math.isclose(audio_cfg_scale, 1.0): + base_params['audio_proj'] = None noise_pred_uncond, teacache_state_uncond = transformer( [z_neg], 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, @@ -2858,6 +2883,24 @@ class WanVideoSampler: noise_pred = noise_pred_uncond + phantom_cfg_scale * (noise_pred_phantom - noise_pred_uncond) + cfg_scale * (noise_pred_cond - noise_pred_phantom) return noise_pred, [teacache_state_cond, teacache_state_uncond, teacache_state_phantom] + #fantasytalking + if fantasytalking_embeds is not None: + if not math.isclose(audio_cfg_scale, 1.0): + if len(teacache_state) != 3: + teacache_state.append(None) + base_params['audio_proj'] = None + noise_pred_no_audio, teacache_state_audio = transformer( + [z_pos], 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, + pred_id=teacache_state[0] if teacache_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 * (noise_pred_cond - noise_pred_no_audio) + return noise_pred, [teacache_state_cond, teacache_state_uncond, teacache_state_audio] + #batched else: teacache_state_uncond = None diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index da05ead..291f1c0 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -322,7 +322,7 @@ class WanSelfAttention(nn.Module): class WanT2VCrossAttention(WanSelfAttention): - def forward(self, x, context, context_lens, clip_embed=None): + def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21): r""" Args: x(Tensor): Shape [B, L1, C] @@ -362,7 +362,7 @@ class WanI2VCrossAttention(WanSelfAttention): self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.attention_mode = attention_mode - def forward(self, x, context, context_lens, clip_embed): + def forward(self, x, context, context_lens, clip_embed, audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21): r""" Args: x(Tensor): Shape [B, L1, C] @@ -389,6 +389,26 @@ class WanI2VCrossAttention(WanSelfAttention): if clip_embed is not None: img_x = img_x.flatten(2) x = x + img_x + + # FantasyTalking audio attention + if audio_proj is not None: + if len(audio_proj.shape) == 4: + audio_q = q.view(b * num_latent_frames, -1, n, d) # [b, 21, l1, n, d] + ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d) + ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d) + audio_x = attention( + audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode + ) + audio_x = audio_x.view(b, q.size(1), n, d) + audio_x = audio_x.flatten(2) + elif len(audio_proj.shape) == 3: + ip_key = self.k_proj(audio_proj).view(b, -1, n, d) + ip_value = self.v_proj(audio_proj).view(b, -1, n, d) + audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode) + audio_x = audio_x.flatten(2) + + x = x + audio_x * audio_scale + x = self.o(x) return x @@ -469,7 +489,11 @@ class WanAttentionBlock(nn.Module): video_attention_split_steps=[], rope_func = "default", clip_embed=None, - camera_embed=None + camera_embed=None, + audio_proj=None, + audio_context_lens=None, + audio_scale=1.0, + num_latent_frames=21, ): r""" @@ -522,12 +546,15 @@ class WanAttentionBlock(nn.Module): if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1: x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) else: - x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) + 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) del e return x @torch.compiler.disable() - def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None): - x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed) + 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): + 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) y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) x = x + (y * e[5]) return x @@ -1023,6 +1050,10 @@ class WanModel(ModelMixin, ConfigMixin): fps_embeds=None, fun_ref = None, fun_camera=None, + audio_proj=None, + audio_context_lens=None, + audio_scale=1.0, + ): r""" Forward pass through the diffusion model @@ -1240,6 +1271,10 @@ class WanModel(ModelMixin, ConfigMixin): current_step=current_step, video_attention_split_steps=self.video_attention_split_steps, camera_embed=camera_embed, + audio_proj=audio_proj, + audio_context_lens=audio_context_lens, + num_latent_frames = F, + audio_scale=audio_scale ) if vace_data is not None: