This commit is contained in:
kijai
2025-12-20 03:14:49 +02:00
parent 93f7af6dc8
commit 4a6e2d3c6c
11 changed files with 713 additions and 135 deletions
+149 -2
View File
@@ -2,6 +2,10 @@ import torch.nn as nn
import torch.nn.functional as F
import torch
import math
from einops import rearrange
from ..wanvideo.modules.model import WanRMSNorm, attention
from ..multitalk.multitalk import RotaryPositionalEmbedding1D, normalize_and_scale
class FeedForwardSwiGLU(nn.Module):
def __init__(
@@ -22,7 +26,7 @@ class FeedForwardSwiGLU(nn.Module):
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
@@ -62,4 +66,147 @@ class TimestepEmbedder(nn.Module):
if t_freq.dtype != dtype:
t_freq = t_freq.to(dtype)
t_emb = self.mlp(t_freq)
return t_emb
return t_emb
class SingleStreamAttention(nn.Module):
def __init__(
self,
dim: int,
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
eps: float = 1e-6,
class_range: int = 24,
class_interval: int = 4,
attention_mode: str = "sdpa",
) -> 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.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.q_norm = WanRMSNorm(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.k_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
# multitalk related params
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 _process_cross_attn(self, x, cond, frames_num=None, x_ref_attn_map=None):
N_t = frames_num
out_dtype = x.dtype
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)) # [B, H, N, D]
q = self.q_norm(q)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
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]
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)
# get kv from encoder_hidden_states
_, N_a, _ = cond.shape
encoder_kv = self.kv_linear(cond)
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)
encoder_k = self.k_norm(encoder_k)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
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)
# Input tensors must be in format ``[B, M, H, K]``, where B is the batch size, M \
# the sequence length, H the number of heads, and K the embeding size per head
q = rearrange(q, "B H M K -> B M H K")
encoder_k = rearrange(encoder_k, "B H M K -> B M H K")
encoder_v = rearrange(encoder_v, "B H M K -> B M H K")
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
x = rearrange(x, "B M H K -> B H M K")
# 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.type(out_dtype)
def forward(self, x, cond, num_latent_frames=None, num_cond_latents=None, x_ref_attn_map=None, human_num=None):
B, N, C = x.shape
if (num_cond_latents is None or num_cond_latents == 0):
# text to video
output = self._process_cross_attn(x, cond, num_latent_frames, x_ref_attn_map)
return None, output
elif num_cond_latents is not None and num_cond_latents > 0:
# image to video or video continuation
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
x_noise = x[:, num_cond_latents_thw:]
cond = rearrange(cond, "(B N_t) M C -> B N_t M C", B=B)
cond = cond[:, num_cond_latents:]
cond = rearrange(cond, "B N_t M C -> (B N_t) M C")
frames_num = num_latent_frames - num_cond_latents
if human_num is not None and human_num == 2:
# multitalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num, x_ref_attn_map)
else:
# singletalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num)
output_cond = torch.zeros((B, num_cond_latents_thw, C), dtype=output_noise.dtype, device=output_noise.device)
return output_cond, output_noise
else:
raise NotImplementedError
+106
View File
@@ -0,0 +1,106 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoLongCatAvatarExtendEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
"audio_embeds": ("MULTITALK_EMBEDS", {"tooltip": "Full length audio embeddings"}),
"num_frames": ("INT", {"default": 93, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
"overlap": ("INT", {"default": 13, "min": 0, "max": 16, "step": 1, "tooltip": "Number of overlapping frames from previous latents" }),
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
"if_not_enough_audio": (["pad_with_start", "mirror_from_end"], {"default": "pad_with_start", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
},
"optional": {
"ref_latent": ("LATENT", {"default": None, "tooltip": "Reference latent for the first frame (used for consistency)"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed=0, ref_latent=None):
new_audio_embed = audio_embeds.copy()
audio_features = torch.stack(new_audio_embed["audio_features"])
print("audio_features shape: ", audio_features.shape)
if audio_features.shape[1] < frames_processed + num_frames:
deficit = frames_processed + num_frames - audio_features.shape[1]
if if_not_enough_audio == "pad_with_start":
pad = audio_features[:, :1].repeat(1, deficit, 1, 1, 1)
audio_features = torch.cat([audio_features, pad], dim=1)
elif if_not_enough_audio == "mirror_from_end":
to_add = audio_features[:, -deficit:, :].flip(dims=[1])
audio_features = torch.cat([audio_features, to_add], dim=1)
log.info(f"Not enough audio features, extended from {new_audio_embed['audio_features'].shape[1]} to {audio_features.shape[1]} frames.")
ref_target_masks = new_audio_embed.get("ref_target_masks", None)
if ref_target_masks is not None:
new_audio_embed["ref_target_masks"] = ref_target_masks[:, frames_processed:frames_processed+num_frames, :]
latent_overlap = (overlap - 1) // 4 + 1
print("prev_latents shape: ", prev_latents["samples"].shape, "latent_overlap: ", latent_overlap)
prev_samples = prev_latents["samples"][:, :, -latent_overlap:].clone()
ref_sample = None
if ref_latent is not None:
ref_sample = ref_latent["samples"][0, :, :1].clone()
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
new_latent_frames = (num_frames - 1) // 4 + 1
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
print("target_shape: ", target_shape)
audio_stride = 2
indices = torch.arange(2 * 2 + 1) - 2
if frames_processed == 0:
audio_start_idx = 0
else:
audio_start_idx = (frames_processed - overlap) * audio_stride
audio_end_idx = audio_start_idx + num_frames * audio_stride
log.info(f"Extracting audio embeddings from index {audio_start_idx} to {audio_end_idx}")
#center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
#center_indices = torch.clamp(center_indices, min=0, max=audio_features.shape[0]-1)
#audio_emb = audio_features[center_indices][None,...]
audio_embs = []
for human_idx in range(len(audio_features)):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_features[human_idx].shape[0] - 1)
audio_emb = audio_features[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_emb = torch.cat(audio_embs, dim=0)
new_audio_embed["audio_features"] = None
new_audio_embed["audio_emb_slice"] = audio_emb
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
"extra_latents": [{"samples": prev_samples, "index": 0}],
"multitalk_embeds": new_audio_embed,
"longcat_ref_latent": ref_sample,
}
return (embeds,)
NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
}
+1
View File
@@ -48,6 +48,7 @@ OPTIONAL_MODULES = [
(".onetoall.nodes", "OneToAll"),
(".WanMove.nodes", "WanMove"),
(".SCAIL.nodes", "SCAIL"),
(".LongCat.nodes", "LongCat"),
]
def register_nodes(module_path: str, name: str, optional: bool) -> None:
+19 -1
View File
@@ -7,6 +7,8 @@ from ..utils import log, set_module_tensor_to_device
import os
import json
import datetime
import scipy.signal as ss
import numpy as np
script_directory = os.path.dirname(os.path.abspath(__file__))
folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2"))
@@ -134,6 +136,15 @@ def loudness_norm(audio_array, sr=16000, lufs=-23):
return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
class MultiTalkWav2VecEmbeds:
@classmethod
@@ -153,6 +164,8 @@ class MultiTalkWav2VecEmbeds:
"audio_3": ("AUDIO",),
"audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space. Supply one mask per speaker (plus optional background) to guide mouth assignment"}),
"add_noise_floor": ("BOOLEAN", {"default": False, "tooltip": "Add a low-level noise floor to the audio to reduce silent gaps"}),
"smooth_transients": ("BOOLEAN", {"default": False, "tooltip": "Apply a low-pass filter to the audio to smooth out transients"}),
}
}
@@ -161,7 +174,8 @@ class MultiTalkWav2VecEmbeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None,
ref_target_masks=None, add_noise_floor=False, smooth_transients=False):
model_type = wav2vec_model["model_type"]
if not "tencent" in model_type.lower():
raise ValueError("Only tencent wav2vec2 models supported by MultiTalk")
@@ -207,6 +221,10 @@ class MultiTalkWav2VecEmbeds:
if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr)
if add_noise_floor:
audio_segment = _add_noise_floor(audio_segment, noise_db=-45)
if smooth_transients:
audio_segment = _smooth_transients(audio_segment, sr=sr)
audio_feature = np.squeeze(
wav2vec2_feature_extractor(audio_segment, sampling_rate=sr).input_values
+1 -1
View File
@@ -2046,7 +2046,7 @@ class WanVideoDecode:
video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0)
return video.cpu().float(),
latents = samples["samples"]
latents = samples["samples"].clone()
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
drop_last = samples.get("drop_last", False)
+33
View File
@@ -1466,6 +1466,39 @@ class WanVideoModelLoader:
sd.update(extra_sd)
del extra_sd
elif "multitalk_audio_proj.proj1.weight" in sd:
log.info("MultiTalk/InfiniteTalk model detected, patching model...")
from .multitalk.multitalk import AudioProjModel
from .wanvideo.modules.model import WanLayerNorm
from .LongCat.layers import SingleStreamAttention
audio_window = 5
vae_scale = 4
for block in transformer.blocks:
with init_empty_weights():
if "blocks.0.audio_modulation.1.weight" in sd:
block.audio_modulation = nn.Sequential(nn.SiLU(), nn.Linear(512, 3 * dim, bias=True))
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
qk_norm=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
multitalk_proj_model = AudioProjModel(
seq_len=audio_window,
seq_len_vf=audio_window+vae_scale-1,
intermediate_dim=512,
output_dim=768,
context_tokens=32,
norm_output_audio=True,
)
transformer.multitalk_audio_proj = multitalk_proj_model
# FlashVSR
if "LQ_proj_in.norm1.gamma" in sd:
+101 -97
View File
@@ -294,10 +294,11 @@ class WanVideoSampler:
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = image_cond_neg =None
vace_data = vace_context = vace_scale = None
fun_or_fl2v_model = has_ref = drop_last = False
fun_or_fl2v_model = drop_last = False
phantom_latents = fun_ref_image = ATI_tracks = None
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
humo_audio = humo_audio_neg = None
has_ref = image_embeds.get("has_ref", False)
#I2V
image_cond = image_embeds.get("image_embeds", None)
@@ -363,15 +364,11 @@ class WanVideoSampler:
control_camera_end_percent = control_embeds.get("control_camera_end_percent", 1.0)
drop_last = image_embeds.get("drop_last", False)
has_ref = image_embeds.get("has_ref", False)
else: #t2v
target_shape = image_embeds.get("target_shape", None)
if target_shape is None:
raise ValueError("Empty image embeds must be provided for T2V models")
has_ref = image_embeds.get("has_ref", False)
# VACE
vace_context = image_embeds.get("vace_context", None)
vace_scale = image_embeds.get("vace_scale", None)
@@ -633,27 +630,34 @@ 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}")
elif multitalk_embeds is not None:
# MultiTalk
multitalk_audio_embeds = audio_emb_slice = audio_features_in = None
multitalk_embeds = image_embeds.get("multitalk_embeds", multitalk_embeds)
if multitalk_embeds is not None:
audio_emb_slice = multitalk_embeds.get("audio_emb_slice", None) # if already sliced
print("audio_emb_slice:", audio_emb_slice.shape)
# Handle single or multiple speaker embeddings
audio_features_in = multitalk_embeds.get("audio_features", None)
if audio_features_in is None:
multitalk_audio_embeds = None
else:
if audio_emb_slice is None:
audio_features_in = multitalk_embeds.get("audio_features", None)
if audio_features_in is not None:
if isinstance(audio_features_in, list):
multitalk_audio_embeds = [emb.to(device, dtype) for emb in audio_features_in]
else:
# keep backward-compatibility with single tensor input
multitalk_audio_embeds = [audio_features_in.to(device, dtype)]
shapes = [tuple(e.shape) for e in multitalk_audio_embeds]
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
audio_scale = multitalk_embeds.get("audio_scale", 1.0)
audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0)
ref_target_masks = multitalk_embeds.get("ref_target_masks", None)
if not isinstance(audio_cfg_scale, list):
audio_cfg_scale = [audio_cfg_scale] * (steps + 1)
shapes = [tuple(e.shape) for e in multitalk_audio_embeds]
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
# FantasyPortrait
fantasy_portrait_input = None
fantasy_portrait_embeds = image_embeds.get("portrait_embeds", None)
@@ -822,7 +826,7 @@ class WanVideoSampler:
# extra latents (Pusa) and 5b
latents_to_insert = add_index = noise_multipliers = None
extra_latents = image_embeds.get("extra_latents", None)
all_indices = []
clean_latent_indices = []
noise_multiplier_list = image_embeds.get("pusa_noise_multipliers", None)
if noise_multiplier_list is not None:
if len(noise_multiplier_list) != latent_video_length:
@@ -832,7 +836,7 @@ class WanVideoSampler:
log.info(f"Using Pusa noise multipliers: {noise_multipliers}")
if extra_latents is not None and transformer.multitalk_model_type.lower() != "infinitetalk":
if noise_multiplier_list is not None:
noise_multiplier_list = list(noise_multiplier_list) + [1.0] * (len(all_indices) - len(noise_multiplier_list))
noise_multiplier_list = list(noise_multiplier_list) + [1.0] * (len(clean_latent_indices) - len(noise_multiplier_list))
for i, entry in enumerate(extra_latents):
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
@@ -843,9 +847,9 @@ class WanVideoSampler:
if start_step == 0:
noise[:, add_index:add_index+num_extra_frames] = entry["samples"].to(noise)
log.info(f"Adding extra samples to latent indices {add_index} to {add_index+num_extra_frames-1}")
all_indices.extend(range(add_index, add_index+num_extra_frames))
clean_latent_indices.extend(range(add_index, add_index+num_extra_frames))
if noise_multipliers is not None and len(noise_multiplier_list) != latent_video_length:
for i, idx in enumerate(all_indices):
for i, idx in enumerate(clean_latent_indices):
noise_multipliers[idx] = noise_multiplier_list[i]
log.info(f"Using Pusa noise multipliers: {noise_multipliers}")
@@ -871,6 +875,17 @@ class WanVideoSampler:
latent = noise
# LongCat-Avatar
longcat_ref_latent = image_embeds.get("longcat_ref_latent", None)
if longcat_ref_latent is not None:
latent = torch.cat([longcat_ref_latent.to(latent), latent], dim=1)
seq_len = math.ceil((latent.shape[2] * latent.shape[3]) / 4 * latent.shape[1])
insert_len = longcat_ref_latent.shape[1]
clean_latent_indices = list(range(0, insert_len)) + [i + insert_len for i in clean_latent_indices]
latent_video_length += insert_len
print("clean_latent_indices:", clean_latent_indices)
audio_stride = 2 if transformer.is_longcat else 1
#controlnet
controlnet_latents = controlnet = None
if transformer_options is not None:
@@ -1387,27 +1402,31 @@ class WanVideoSampler:
else:
z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
if not multitalk_sampling and multitalk_audio_embeds is not None:
multitalk_audio_input = None
if audio_emb_slice is not None:
print("audio_emb_slice shape: ", audio_emb_slice.shape)
multitalk_audio_input = audio_emb_slice.to(z)
elif not multitalk_sampling and multitalk_audio_embeds is not None:
audio_embedding = multitalk_audio_embeds
audio_embs = []
indices = (torch.arange(4 + 1) - 2) * 1
human_num = len(audio_embedding)
# split audio with window size
audio_end_idx = latent_video_length * 4 + 1 if add_cond is not None else (latent_video_length-1) * 4 + 1
audio_end_idx = audio_end_idx * audio_stride
if context_window is None:
for human_idx in range(human_num):
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.arange(0, audio_end_idx, audio_stride).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].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
else:
for human_idx in range(human_num):
audio_start = context_window[0] * 4
audio_end = context_window[-1] * 4 + 1
audio_start = (context_window[0] * 4) * audio_stride
audio_end = (context_window[-1] * 4 + 1) * audio_stride
#print("audio_start: ", audio_start, "audio_end: ", audio_end)
center_indices = torch.arange(audio_start, audio_end, 1).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.arange(audio_start, audio_end, audio_stride).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].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
@@ -1515,7 +1534,7 @@ class WanVideoSampler:
"add_cond": add_cond_input, # additional conditioning input
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
"multitalk_audio": multitalk_audio_input if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk audio input
"multitalk_audio": multitalk_audio_input, # Multi/InfiniteTalk audio input
"ref_target_masks": ref_target_masks if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk reference target masks
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
"standin_input": standin_input, # Stand-in reference input
@@ -1545,7 +1564,8 @@ class WanVideoSampler:
"ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi
"flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
"longcat_num_cond_latents": len(clean_latent_indices) if transformer.is_longcat else 0,
"longcat_num_ref_latents": longcat_ref_latent.shape[1] if longcat_ref_latent is not None else 0,
"sdancer_input": sdancer_input, # SteadyDancer input
"one_to_all_input": one_to_all_data, # One-to-All input
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
@@ -1577,7 +1597,16 @@ class WanVideoSampler:
if math.isclose(cfg_scale, 1.0):
if use_fresca:
noise_pred_cond = fourier_filter(noise_pred_cond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
return noise_pred_cond, noise_pred_ovi, [cache_state_cond]
if multitalk_audio_input is not None and not math.isclose(audio_cfg_scale[idx], 1.0):
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
noise_pred_uncond_audio, _, cache_state_uncond = transformer(
context=positive_embeds, pred_id=cache_state[0] if cache_state else None,
vace_data=vace_data, attn_cond=attn_cond, **base_params)
return noise_pred_uncond_audio[0] + audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_uncond_audio[0]), noise_pred_ovi, [cache_state_cond, cache_state_uncond]
else:
return noise_pred_cond, noise_pred_ovi, [cache_state_cond]
#unconditional (negative) pass
base_params['is_uncond'] = True
@@ -1591,12 +1620,12 @@ class WanVideoSampler:
if neg_latent is not None:
base_params['x'] = [torch.cat([z[:, :-humo_reference_count], neg_latent], dim=1)]
noise_pred_uncond, noise_pred_ovi_uncond, cache_state_uncond = transformer(
noise_pred_uncond_text, noise_pred_ovi_uncond, cache_state_uncond = transformer(
context=negative_embeds if humo_audio_input_neg is None else positive_embeds, #ti #t
pred_id=cache_state[1] if cache_state else None,
vace_data=vace_data, attn_cond=attn_cond_neg,
**base_params)
noise_pred_uncond = noise_pred_uncond[0]
noise_pred_uncond_text = noise_pred_uncond_text[0]
noise_pred_ovi_uncond = noise_pred_ovi_uncond[0] if noise_pred_ovi_uncond is not None else None
# HuMo
@@ -1611,8 +1640,8 @@ class WanVideoSampler:
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + humo_audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_humo_audio_uncond[0])
+ (cfg_scale - 2.0) * (noise_pred_humo_audio_uncond[0] - noise_pred_uncond))
noise_pred = (noise_pred_uncond_text + humo_audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_humo_audio_uncond[0])
+ (cfg_scale - 2.0) * (noise_pred_humo_audio_uncond[0] - noise_pred_uncond_text))
return noise_pred, None, [cache_state_cond, cache_state_uncond, cache_state_humo]
elif humo_audio_input is not None:
if cache_state is not None and len(cache_state) != 4:
@@ -1628,8 +1657,8 @@ class WanVideoSampler:
context=positive_embeds, pred_id=cache_state[3] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (humo_audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_humo_audio[0])
+ cfg_scale * (noise_pred_humo_audio[0] - noise_pred_uncond)
+ cfg_scale * (noise_pred_uncond - noise_pred_humo_null[0])
+ cfg_scale * (noise_pred_humo_audio[0] - noise_pred_uncond_text)
+ cfg_scale * (noise_pred_uncond_text - noise_pred_humo_null[0])
+ noise_pred_humo_null[0])
return noise_pred, None, [cache_state_cond, cache_state_uncond, cache_state_humo, cache_state_humo2]
@@ -1641,32 +1670,28 @@ class WanVideoSampler:
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + phantom_cfg_scale[idx] * (noise_pred_phantom[0] - noise_pred_uncond)
noise_pred = (noise_pred_uncond_text + phantom_cfg_scale[idx] * (noise_pred_phantom[0] - noise_pred_uncond_text)
+ cfg_scale * (noise_pred_cond - noise_pred_phantom[0]))
return noise_pred, None,[cache_state_cond, cache_state_uncond, cache_state_phantom]
# audio cfg (fantasytalking and multitalk)
if (fantasytalking_embeds is not None or multitalk_audio_embeds is not None):
if (fantasytalking_embeds is not None or multitalk_audio_input 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)
# Set audio parameters to None/zeros based on type
if fantasytalking_embeds is not None:
base_params['audio_proj'] = None
audio_context = positive_embeds
else: # multitalk
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
audio_context = negative_embeds
base_params['audio_proj'] = None
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:] if multitalk_audio_input is not None else None
base_params['is_uncond'] = False
noise_pred_no_audio, _, cache_state_audio = transformer(
context=audio_context,
noise_pred_uncond_audio, _, cache_state_audio = transformer(
context=negative_embeds,
pred_id=cache_state[2] if cache_state else None,
vace_data=vace_data,
**base_params)
noise_pred_uncond_audio = noise_pred_uncond_audio[0]
noise_pred = (noise_pred_uncond
+ cfg_scale * (noise_pred_no_audio[0] - noise_pred_uncond)
+ audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_no_audio[0]))
noise_pred = noise_pred_uncond_audio + cfg_scale * (
(noise_pred_cond - noise_pred_uncond_text)
+ audio_cfg_scale[idx] * (noise_pred_uncond_text - noise_pred_uncond_audio))
return noise_pred, None,[cache_state_cond, cache_state_uncond, cache_state_audio]
# lynx
if lynx_embeds is not None and not math.isclose(lynx_cfg_scale[idx], 1.0):
@@ -1677,7 +1702,7 @@ class WanVideoSampler:
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + lynx_cfg_scale[idx] * (noise_pred_lynx[0] - noise_pred_uncond)
noise_pred = (noise_pred_uncond_text + lynx_cfg_scale[idx] * (noise_pred_lynx[0] - noise_pred_uncond_text)
+ cfg_scale * (noise_pred_cond - noise_pred_lynx[0]))
return noise_pred, None, [cache_state_cond, cache_state_uncond, cache_state_lynx]
# one-to-all
@@ -1691,7 +1716,7 @@ class WanVideoSampler:
context=negative_embeds, pred_id=cache_state[2] if cache_state else None, vace_data=None,
**base_params)
noise_pred = (noise_pred_uncond + one_to_all_pose_cfg_scale[idx] * (noise_pred_pose_uncond[0] - noise_pred_uncond)
noise_pred = (noise_pred_uncond_text + one_to_all_pose_cfg_scale[idx] * (noise_pred_pose_uncond[0] - noise_pred_uncond_text)
+ cfg_scale * (noise_pred_cond - noise_pred_pose_uncond[0]))
return noise_pred, None, [cache_state_cond, cache_state_uncond, cache_state_ref]
@@ -1721,23 +1746,23 @@ class WanVideoSampler:
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
noise_pred_uncond_scaled = noise_pred_uncond * alpha
noise_pred_uncond_text = noise_pred_uncond_text * alpha
if use_tangential:
noise_pred_uncond_scaled = tangential_projection(noise_pred_cond, noise_pred_uncond_scaled)
noise_pred_uncond_text = tangential_projection(noise_pred_cond, noise_pred_uncond_text)
# RAAG (RATIO-aware Adaptive Guidance)
if raag_alpha > 0.0:
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_scaled, cfg_scale, raag_alpha)
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_text, cfg_scale, raag_alpha)
log.info(f"RAAG modified cfg: {cfg_scale}")
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(noise_pred_cond - noise_pred_uncond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
noise_pred = noise_pred_uncond_text + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
noise_pred = noise_pred_uncond_text + cfg_scale * (noise_pred_cond - noise_pred_uncond_text)
del noise_pred_uncond_text, noise_pred_cond
if latent_model_input_ovi is not None:
if ovi_audio_cfg is None:
@@ -1780,7 +1805,6 @@ class WanVideoSampler:
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
#torch.cuda.memory._record_memory_history(max_entries=100000)
except:
pass
@@ -1836,7 +1860,7 @@ class WanVideoSampler:
# Set latent for denoising
latent = current_latent
if is_pusa and all_indices:
if is_pusa and clean_latent_indices:
pusa_noisy_steps = image_embeds.get("pusa_noisy_steps", -1)
if pusa_noisy_steps == -1:
pusa_noisy_steps = len(timesteps)
@@ -1867,15 +1891,15 @@ class WanVideoSampler:
current_step_percentage = idx / len(timesteps)
timestep = torch.tensor([t]).to(device)
if is_pusa or ((is_5b or transformer.is_longcat) and all_indices):
if is_pusa or ((is_5b or transformer.is_longcat) and clean_latent_indices):
orig_timestep = timestep
timestep = timestep.unsqueeze(1).repeat(1, latent_video_length)
if extra_latents is not None:
if all_indices and noise_multipliers is not None:
if clean_latent_indices and noise_multipliers is not None:
if is_pusa:
scheduler_step_args["cond_frame_latent_indices"] = all_indices
scheduler_step_args["cond_frame_latent_indices"] = clean_latent_indices
scheduler_step_args["noise_multipliers"] = noise_multipliers
for latent_idx in all_indices:
for latent_idx in clean_latent_indices:
timestep[:, latent_idx] = timestep[:, latent_idx] * noise_multipliers[latent_idx]
# add noise for conditioning frames if multiplier > 0
if idx < pusa_noisy_steps and noise_multipliers[latent_idx] > 0:
@@ -1889,7 +1913,7 @@ class WanVideoSampler:
timestep_cond[:, latent_idx:latent_idx+1].to(device),
noise_multiplier=noise_multipliers[latent_idx])
else:
timestep[:, all_indices] = 0
timestep[:, clean_latent_indices] = 0
#print("timestep: ", timestep)
### latent shift
@@ -2259,8 +2283,8 @@ class WanVideoSampler:
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = step_iteration_count= 0
audio_end_idx = audio_start_idx + clip_length
audio_start_idx = iteration_count = step_iteration_count = 0
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
@@ -2309,7 +2333,7 @@ class WanVideoSampler:
audio_embs = []
# split audio with window size
for human_idx in range(human_num):
center_indices = torch.arange(audio_start_idx, audio_end_idx, 1).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).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].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
@@ -3135,28 +3159,10 @@ class WanVideoSampler:
if transformer.is_longcat:
noise_pred = -noise_pred
if len(timestep.shape) != 1 and not is_pusa: #5b and longcat
# all_indices is a list of indices to skip
total_indices = list(range(latent.shape[1]))
process_indices = [i for i in total_indices if i not in all_indices]
if process_indices:
latent_to_process = latent[:, process_indices]
noise_pred_to_process = noise_pred[:, process_indices]
latent_slice = sample_scheduler.step(
noise_pred_to_process.unsqueeze(0),
orig_timestep,
latent_to_process.unsqueeze(0),
**scheduler_step_args
)[0].squeeze(0)
# Reconstruct the latent tensor: keep skipped indices as-is, update others
new_latent = []
for i in total_indices:
if i in all_indices:
new_latent.append(latent[:, i:i+1])
else:
j = process_indices.index(i)
new_latent.append(latent_slice[:, j:j+1])
latent = torch.cat(new_latent, dim=1)
if len(timestep.shape) != 1 and clean_latent_indices and not is_pusa: #5b and longcat, skip clean latents for scheduler step
step_process_indices = [i for i in range(latent.shape[1]) if i not in clean_latent_indices]
latent[:, step_process_indices] = sample_scheduler.step(noise_pred[:, step_process_indices].unsqueeze(0), orig_timestep,
latent[:, step_process_indices].unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
else:
if latents_to_not_step > 0:
raw_latent = latent[:, :latents_to_not_step]
@@ -3169,11 +3175,7 @@ class WanVideoSampler:
noise_pred_in = noise_pred
latent = sample_scheduler.step(noise_pred_in.unsqueeze(0), timestep, latent.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
if noise_pred_flipped is not None:
latent_backwards = sample_scheduler_flipped.step(
noise_pred_flipped.unsqueeze(0),
timestep,
latent_flipped.unsqueeze(0),
**scheduler_step_args)[0].squeeze(0)
latent_backwards = sample_scheduler_flipped.step(noise_pred_flipped.unsqueeze(0), timestep, latent_flipped.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
if latents_to_not_step > 0:
@@ -3243,6 +3245,8 @@ class WanVideoSampler:
latent = latent[:,:-phantom_latents.shape[1]]
if humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
if longcat_ref_latent is not None:
latent = latent[:, longcat_ref_latent.shape[1]:]
cache_states = None
if cache_args is not None:
@@ -3261,8 +3265,6 @@ class WanVideoSampler:
try:
print_memory(device)
#torch.cuda.memory._dump_snapshot("wanvideowrapper_memory_dump.pt")
#torch.cuda.memory._record_memory_history(enabled=None)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
@@ -3381,6 +3383,7 @@ class WanVideoScheduler:
},
"optional": {
"sigmas": ("SIGMAS", ),
"enhance_hf": ("BOOLEAN", {"default": False, "tooltip": "Enhanced high-frequency denoising schedule"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
@@ -3393,9 +3396,9 @@ class WanVideoScheduler:
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None):
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None, enhance_hf=False):
sample_scheduler, timesteps, start_idx, end_idx = get_scheduler(
scheduler, steps, start_step, end_step, shift, device, sigmas=sigmas, log_timesteps=True)
scheduler, steps, start_step, end_step, shift, device, sigmas=sigmas, log_timesteps=True, enhance_hf=enhance_hf)
scheduler_dict = {
"sample_scheduler": sample_scheduler,
@@ -3480,6 +3483,7 @@ class WanVideoSchedulerv2(WanVideoScheduler):
},
"optional": {
"sigmas": ("SIGMAS", ),
"enhance_hf": ("BOOLEAN", {"default": False, "tooltip": "Enhanced high-frequency denoising schedule"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.4.3"
version = "1.4.4"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+126 -32
View File
@@ -633,15 +633,15 @@ class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
inner_t=None, inner_c=None, cross_freqs=None,
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, num_cond_latents=None, **kwargs):
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
s = x.size(1)
# compute query
is_longcat = x.shape[-1] == 4096
if is_longcat:
if num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (s // num_latent_frames)
if longcat_num_cond_latents is not None and longcat_num_cond_latents > 0:
num_cond_latents_thw = longcat_num_cond_latents * (s // num_latent_frames)
x = x[:, num_cond_latents_thw:]
q = self.norm_q(self.q(x).view(b, -1, n, d))
else:
@@ -712,7 +712,7 @@ class WanT2VCrossAttention(WanSelfAttention):
x = x.add(target_x)
if is_longcat and num_cond_latents is not None and num_cond_latents > 0:
if is_longcat and longcat_num_cond_latents > 0:
return torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), self.o(x)], dim=1).contiguous()
return self.o(x)
@@ -914,7 +914,7 @@ class WanAttentionBlock(nn.Module):
from ...LongCat.layers import FeedForwardSwiGLU
mlp_ratio = 4
self.ffn = FeedForwardSwiGLU(dim=self.dim, hidden_dim=int(self.dim * mlp_ratio))
# modulation
if not is_longcat:
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
@@ -1003,7 +1003,7 @@ class WanAttentionBlock(nn.Module):
humo_audio_input=None, humo_audio_scale=1.0, #humo audio
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None,
num_cond_latents=None, #longcat image cond amount
longcat_num_cond_latents=0, #longcat image cond amount
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
e_tr=None, tr_num=0, tr_start=0, #token replacement
):
@@ -1030,7 +1030,7 @@ class WanAttentionBlock(nn.Module):
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
#del e
input_dtype = x.dtype
B, N, C = x.shape
T = num_latent_frames
@@ -1181,18 +1181,64 @@ class WanAttentionBlock(nn.Module):
full_k = torch.cat([k, k_ip], dim=1)
full_v = torch.cat([v, v_ip], dim=1)
y = self.self_attn.forward(q, full_k, full_v, seq_lens)
elif is_longcat and num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
# process the condition tokens
x_cond = self.self_attn.forward(
q[:, :num_cond_latents_thw].contiguous(),
k[:, :num_cond_latents_thw].contiguous(),
v[:, :num_cond_latents_thw].contiguous(),
seq_lens)
# process the noise tokens
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif is_longcat and longcat_num_cond_latents > 0:
if longcat_num_cond_latents == 1:
num_cond_latents_thw = longcat_num_cond_latents * (N // num_latent_frames)
# process the condition tokens
x_cond = self.self_attn.forward(
q[:, :num_cond_latents_thw].contiguous(),
k[:, :num_cond_latents_thw].contiguous(),
v[:, :num_cond_latents_thw].contiguous(),
seq_lens)
# process the noise tokens
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif longcat_num_cond_latents > 1: # video continuation
num_ref_latents_thw = (N // num_latent_frames)
num_cond_latents_thw = longcat_num_cond_latents * (N // num_latent_frames)
# process the condition tokens
q_ref = q[:, :num_ref_latents_thw].contiguous()
k_ref = k[:, :num_ref_latents_thw].contiguous()
v_ref = v[:, :num_ref_latents_thw].contiguous()
q_cond = q[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
k_cond = k[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
v_cond = v[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens)
x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens)
if longcat_num_cond_latents == num_latent_frames:
y = torch.cat([x_ref, x_cond], dim=1).contiguous()
else:
# process the noise tokens
q_noise = q[:, num_cond_latents_thw:].contiguous()
start_noise, end_noise, num_noisy_frames = 0, 0, num_latent_frames - longcat_num_cond_latents
mask_frame_range = 3 #todo: make it configurable?
ref_img_index = 10 #todo: make it configurable?
num_ref_latents = 1 # todo: make it configurable?
if mask_frame_range is not None and mask_frame_range > 0:
start_noise = ref_img_index - mask_frame_range - longcat_num_cond_latents + num_ref_latents
end_noise = ref_img_index + mask_frame_range - longcat_num_cond_latents + num_ref_latents + 1
if start_noise >= 0 and end_noise > start_noise and end_noise <= num_noisy_frames:
# remove attention with the reference image in the target range, preventing repeated actions.
start_pos = start_noise * (N // num_latent_frames)
end_pos = end_noise * (N // num_latent_frames)
q_noise_front = q_noise[:, :start_pos].contiguous()
q_noise_maskref = q_noise[:, start_pos:end_pos].contiguous()
q_noise_back = q_noise[:, end_pos:].contiguous()
k_non_ref = k[:, num_ref_latents_thw:].contiguous()
v_non_ref = v[:, num_ref_latents_thw:].contiguous()
x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens) # q_front has attention with ref + cond + noisy
x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens) # q_back has attention with ref + cond + noisy
x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens) # q_mask has attention with cond+noisy
x_noise = torch.cat([x_noise_front, x_noise_maskref, x_noise_back], dim=1).contiguous()
else:
x_noise = self.self_attn.forward(q_noise, k, v, seq_lens)
# merge x_cond and x_noise
y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous()
else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale)
@@ -1263,12 +1309,23 @@ class WanAttentionBlock(nn.Module):
x = x + self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context,
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents)
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, longcat_num_cond_latents=longcat_num_cond_latents)
x = x.to(input_dtype)
# MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
if is_longcat:
audio_output_cond, x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), multitalk_audio_embedding, num_latent_frames=num_latent_frames,
num_cond_latents=longcat_num_cond_latents, x_ref_attn_map=x_ref_attn_map, human_num=human_num)
audio_shift_mca, audio_scale_mca, audio_gate_mca = self.audio_modulation(e[:, longcat_num_cond_latents:]).unsqueeze(2).chunk(3, dim=-1) # [B, T, 1, C]
x_audio = self.modulate(self.norm1(x_audio.view(B, T-longcat_num_cond_latents, -1, C).to(audio_shift_mca.dtype)), audio_shift_mca, audio_scale_mca, seg_idx=self.seg_idx).to(input_dtype).view(B, -1, C)
x_audio = (x_audio.view(B, T-longcat_num_cond_latents, -1, C).float() * audio_gate_mca).to(input_dtype).view(B, -1, C)
if audio_output_cond is not None:
x_audio = torch.cat([audio_output_cond, x_audio], dim=1).contiguous()
else:
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
x = x.add(x_audio, alpha=audio_scale)
# MTV-Crafter Motion Attention
@@ -1282,7 +1339,7 @@ class WanAttentionBlock(nn.Module):
# ffn
if self.rope_func == "comfy_chunked":
if self.rope_func == "comfy_chunked" and not is_longcat and not use_token_replace and not zero_timestep:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
x_ffn = self.ffn_chunked(mod_x)
else:
@@ -2128,7 +2185,8 @@ class WanModel(torch.nn.Module):
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None,
steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None,
ref_frame_index=10, longcat_num_ref_latents=None):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
@@ -2144,7 +2202,19 @@ class WanModel(torch.nn.Module):
# Main frames position IDs
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
if longcat_num_ref_latents > 0:
# Create temporal grid with ref_frame_index prepended, followed by sequential frames
grid_t = torch.cat([
torch.tensor([ref_frame_index], dtype=dtype, device=device),
torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device)
], dim=0)
print("grid_t:", grid_t)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1)
else:
# Standard temporal encoding
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
@@ -2243,7 +2313,7 @@ class WanModel(torch.nn.Module):
lynx_embeds=None,
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None,
longcat_num_cond_latents=0, longcat_num_ref_latents=0, # for LongCat
add_text_emb=None,
sdancer_input=None, # SteadyDancer
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
@@ -2331,6 +2401,7 @@ class WanModel(torch.nn.Module):
freqs = freqs.to(device)
_, F, H, W = x[0].shape
print("Input shape:", x[0].shape)
ref_frame_shape = pose_frame_shape = None
sdancer_enabled = False
@@ -2558,6 +2629,7 @@ class WanModel(torch.nn.Module):
tuple(pose_frame_shape) if pose_frame_shape is not None else None,
self.rope_embedder.k,
tuple(ntk_alphas),
longcat_num_ref_latents,
)
# Check cache using key comparison
@@ -2573,6 +2645,7 @@ class WanModel(torch.nn.Module):
ntk_alphas=ntk_alphas,
ref_frame_shape=ref_frame_shape,
pose_frame_shape=pose_frame_shape,
longcat_num_ref_latents=longcat_num_ref_latents,
device=x.device,
dtype=x.dtype
)
@@ -2636,14 +2709,19 @@ class WanModel(torch.nn.Module):
e_token_replace = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t_token_replace.flatten()).to(time_embed_dtype)) # b, dim
e0_token_replace = self.time_projection(e_token_replace).unflatten(1, (6, self.dim)) # b, 6, dim
else:
print("input t shape:", t.shape)
print("F:", F)
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
time_embed_dtype = self.base_dtype
if len(t.shape) == 1:
t = t.unsqueeze(1).expand(-1, F) # [B, T]
print("t expanded shape:", t.shape)
self.time_embedding.to(torch.float32)
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32).reshape(1, F, -1)
print("t float shape:", t.float().flatten().shape)
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32)#.reshape(1, F, -1)
print("e0 shape:", e0.shape)
e = e0 = e0.reshape(1, F, -1)
if self.audio_model is not None:
#if t.dim() == 1:
@@ -2778,10 +2856,26 @@ class WanModel(torch.nn.Module):
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.multitalk_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(self.base_dtype)
multitalk_audio_embedding = self.multitalk_audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
self.multitalk_audio_proj.to(self.offload_device)
human_num = len(multitalk_audio_embedding)
# LongCat-Avatar specific
print("longcat_num_cond_latents:", longcat_num_cond_latents, "longcat_num_ref_latents:", longcat_num_ref_latents)
if longcat_num_ref_latents > 0:
audio_start_ref = multitalk_audio_embedding[:, [0], :, :] # padding
multitalk_audio_embedding = torch.cat([audio_start_ref, multitalk_audio_embedding], dim=1).contiguous()
if longcat_num_cond_latents > 0:
multitalk_audio_embedding = multitalk_audio_embedding[:, (-F // self.patch_size[0]):]
if ref_target_masks is not None:
multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(self.base_dtype)
multitalk_audio_embedding = multitalk_audio_embedding.squeeze(0)
else:
multitalk_audio_embedding = rearrange(multitalk_audio_embedding, "b t n c -> (b t) n c")
# convert ref_target_masks to token_ref_target_masks
token_ref_target_masks = None
@@ -2975,7 +3069,7 @@ class WanModel(torch.nn.Module):
lynx_x_ip=lynx_x_ip,
lynx_ip_scale=lynx_ip_scale,
lynx_ref_scale=lynx_ref_scale,
num_cond_latents=num_cond_latents,
longcat_num_cond_latents=longcat_num_cond_latents,
onetoall_ref_scale=onetoall_ref_scale,
e_tr=e0_token_replace if use_token_replace else None,
tr_start=token_replace_start,
+22 -1
View File
@@ -1,9 +1,11 @@
import torch
import numpy as np
from .fm_solvers import (FlowDPMSolverMultistepScheduler)
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
from .basic_flowmatch import FlowMatchScheduler
from .flowmatch_pusa import FlowMatchSchedulerPusa
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
from .ersde_scheduler import ERSDEScheduler
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
from .fm_sa_ode import FlowMatchSAODEStableScheduler
from .fm_rcm import rCMFlowMatchScheduler
@@ -25,6 +27,7 @@ scheduler_list = [
"deis",
"lcm", "lcm/beta",
"res_multistep",
"er_sde",
"flowmatch_causvid",
"flowmatch_distill",
"flowmatch_pusa",
@@ -39,7 +42,7 @@ def _apply_custom_sigmas(sample_scheduler, sigmas, device):
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, **kwargs):
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
timesteps = None
if sigmas is not None:
steps = len(sigmas) - 1
@@ -136,6 +139,12 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'er_sde':
sample_scheduler = ERSDEScheduler(shift=shift)
if sigmas is None:
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif "sa_ode_stable" in scheduler:
sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs)
if sigmas is None:
@@ -152,6 +161,18 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if timesteps is None:
timesteps = sample_scheduler.timesteps
if enhance_hf:
num_tail_uniform_steps = max(3, min(15, int(len(timesteps) * 0.2))) # Use 20% of steps for uniform tail (minimum 3, maximum 15)
tail_uniform_start = float(timesteps.max()) * 0.5 # Split at 50% of the timestep range
tail_uniform_end = 0
timesteps_uniform_tail = list(np.linspace(tail_uniform_start, tail_uniform_end, num_tail_uniform_steps, dtype=np.float32, endpoint=(tail_uniform_end != 0)))
timesteps_uniform_tail = [torch.tensor(t, device=device).unsqueeze(0) for t in timesteps_uniform_tail]
filtered_timesteps = [timestep.unsqueeze(0).to(device) for timestep in timesteps if timestep > tail_uniform_start]
timesteps = torch.cat(filtered_timesteps + timesteps_uniform_tail)
sample_scheduler.timesteps = timesteps
sample_scheduler.sigmas = torch.cat([timesteps / 1000, torch.zeros(1, device=timesteps.device)])
steps = len(timesteps)
if (isinstance(start_step, int) and end_step != -1 and start_step >= end_step) or (not isinstance(start_step, int) and start_step != -1 and end_step >= start_step):
raise ValueError("start_step must be less than end_step")
+154
View File
@@ -0,0 +1,154 @@
import torch
class ERSDEScheduler():
"""Extended Reverse-Time SDE solver (VP ER-SDE-Solver-3).
Based on: arXiv: https://arxiv.org/abs/2309.06169
Code reference: https://github.com/QinpengCui/ER-SDE-Solver/blob/main/er_sde_solver.py
"""
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0,
sigma_max=1.0, sigma_min=0.003 / 1.002, max_stage=3, s_noise=1.0,
num_integration_points=200):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.max_stage = max_stage
self.s_noise = s_noise
self.num_integration_points = num_integration_points
self.set_timesteps(num_inference_steps)
self.old_denoised = None
self.old_denoised_d = None
self.step_index = 0
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, sigmas=None):
"""Generate the full sigma schedule (from max to min)."""
full_sigmas = torch.linspace(self.sigma_max, self.sigma_min, self.num_train_timesteps)
ss = len(full_sigmas) / num_inference_steps
if sigmas is None:
sigmas = []
for x in range(num_inference_steps):
idx = int(round(x * ss))
sigmas.append(float(full_sigmas[idx]))
sigmas.append(0.0)
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas * self.num_train_timesteps
self.step_index = 0
self.old_denoised = None
self.old_denoised_d = None
def default_er_sde_noise_scaler(self, x):
return x * ((x ** 0.3).exp() + 10.0)
def step(self, model_output, timestep, sample, generator):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
if timestep.ndim == 0:
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
noise_scaler = self.default_er_sde_noise_scaler
# Get current and next sigma
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.zeros_like(sigma)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
er_lambda_s = sigma
er_lambda_t = sigma_next
# Calculate alpha values
alpha_s = sigma / (er_lambda_s + 1e-10)
alpha_t = sigma_next / (er_lambda_t + 1e-10)
r_alpha = alpha_t / (alpha_s + 1e-10)
# Denoised prediction (x_0 estimate)
denoised = sample - sigma * model_output
# Determine which stage to use
stage_used = min(self.max_stage, self.step_index + 1)
if sigma_next == 0 or (sigma_next == 0.0).all():
# Final step - return denoised
x = denoised
else:
r = noise_scaler(er_lambda_t) / (noise_scaler(er_lambda_s) + 1e-10)
# Stage 1: Euler step
x = r_alpha * r * sample + alpha_t * (1 - r) * denoised
if stage_used >= 2 and self.old_denoised is not None:
dt = er_lambda_t - er_lambda_s
lambda_step_size = -dt / self.num_integration_points
# Create integration points
point_indice = torch.arange(0, self.num_integration_points,
dtype=torch.float32, device=sample.device)
lambda_pos = er_lambda_t + point_indice * lambda_step_size
scaled_pos = noise_scaler(lambda_pos)
# Stage 2: Second-order correction
s = torch.sum(1 / (scaled_pos + 1e-10)) * lambda_step_size
# Get previous sigma for derivative calculation
if timestep_id > 0:
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1)
er_lambda_prev = sigma_prev
else:
er_lambda_prev = er_lambda_s
denoised_d = (denoised - self.old_denoised) / ((er_lambda_s - er_lambda_prev) + 1e-10)
x = x + alpha_t * (dt + s * noise_scaler(er_lambda_t)) * denoised_d
if stage_used >= 3 and self.old_denoised_d is not None:
# Stage 3: Third-order correction
s_u = torch.sum((lambda_pos - er_lambda_s) / (scaled_pos + 1e-10)) * lambda_step_size
# Get sigma from two steps ago
if timestep_id > 1:
sigma_prev_prev = self.sigmas[timestep_id - 2].reshape(-1, 1, 1, 1)
er_lambda_prev_prev = sigma_prev_prev
else:
er_lambda_prev_prev = er_lambda_prev
denoised_u = (denoised_d - self.old_denoised_d) / (((er_lambda_s - er_lambda_prev_prev) / 2) + 1e-10)
x = x + alpha_t * ((dt ** 2) / 2 + s_u * noise_scaler(er_lambda_t)) * denoised_u
self.old_denoised_d = denoised_d
# Add stochastic noise
if self.s_noise > 0:
noise_term = (er_lambda_t ** 2 - er_lambda_s ** 2 * r ** 2).sqrt()
noise_term = torch.nan_to_num(noise_term, nan=0.0)
noise = torch.randn(*x.shape, dtype=torch.float32, device=torch.device("cpu"), generator=generator).to(x)
x = x + alpha_t * noise * self.s_noise * noise_term
# Store current denoised for next iteration
self.old_denoised = denoised
self.step_index += 1
return x
def add_noise(self, original_samples, noise, timestep):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)