This commit is contained in:
kijai
2025-06-18 15:41:53 +03:00
parent 058286fc0f
commit 58104b620f
7 changed files with 860 additions and 45 deletions
+4
View File
@@ -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)
+21 -6
View File
@@ -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,}
+374
View File
@@ -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
+149
View File
@@ -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",
}
+141
View File
@@ -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,
)
+106 -25
View File
@@ -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:]
+65 -14
View File
@@ -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: