Support FantasyTalking
This commit is contained in:
@@ -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"]
|
||||
@@ -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
|
||||
)
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user