From 4a6e2d3c6c96dbdeb2342dd2a1829c761469b676 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 20 Dec 2025 03:14:49 +0200 Subject: [PATCH] Init --- LongCat/layers.py | 151 ++++++++++++++++++- LongCat/nodes.py | 106 +++++++++++++ __init__.py | 1 + multitalk/nodes.py | 20 ++- nodes.py | 2 +- nodes_model_loading.py | 33 +++++ nodes_sampler.py | 198 +++++++++++++------------ pyproject.toml | 2 +- wanvideo/modules/model.py | 158 ++++++++++++++++---- wanvideo/schedulers/__init__.py | 23 ++- wanvideo/schedulers/ersde_scheduler.py | 154 +++++++++++++++++++ 11 files changed, 713 insertions(+), 135 deletions(-) create mode 100644 LongCat/nodes.py create mode 100644 wanvideo/schedulers/ersde_scheduler.py diff --git a/LongCat/layers.py b/LongCat/layers.py index df98c01..34744d6 100644 --- a/LongCat/layers.py +++ b/LongCat/layers.py @@ -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 \ No newline at end of file + 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 diff --git a/LongCat/nodes.py b/LongCat/nodes.py new file mode 100644 index 0000000..9effe93 --- /dev/null +++ b/LongCat/nodes.py @@ -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", + } diff --git a/__init__.py b/__init__.py index fef124e..4efb858 100644 --- a/__init__.py +++ b/__init__.py @@ -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: diff --git a/multitalk/nodes.py b/multitalk/nodes.py index fed8fb8..837a763 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -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 diff --git a/nodes.py b/nodes.py index c35f5cd..0c78f44 100644 --- a/nodes.py +++ b/nodes.py @@ -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) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index c0de0f2..8683839 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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: diff --git a/nodes_sampler.py b/nodes_sampler.py index cdbda4c..9a3b168 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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", diff --git a/pyproject.toml b/pyproject.toml index a117b7e..427a9c7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index b2d19f8..27512fb 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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, diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index e39e7e4..abb450f 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -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") diff --git a/wanvideo/schedulers/ersde_scheduler.py b/wanvideo/schedulers/ersde_scheduler.py new file mode 100644 index 0000000..6abd7a7 --- /dev/null +++ b/wanvideo/schedulers/ersde_scheduler.py @@ -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)