Author SHA1 Message Date
kijai fd32b14fdc Clean prints 2025-12-23 02:02:31 +02:00
kijai 1776695e26 Update nodes_model_loading.py 2025-12-23 01:48:12 +02:00
kijai ef36204fa8 Reduce peak VRAM use 2025-12-23 01:35:08 +02:00
kijai c6f32c1424 Norm dtype 2025-12-22 23:53:41 +02:00
kijai 6d4a0f6e53 Merge branch 'main' into longcat_avatar 2025-12-22 22:11:38 +02:00
kijai e7e00061e5 Update nodes_sampler.py 2025-12-22 00:43:01 +02:00
kijai eb5ec262a0 Merge branch 'main' into longcat_avatar 2025-12-22 00:42:53 +02:00
kijai 7c0ba84a26 remove prints 2025-12-21 23:00:43 +02:00
kijai 06a86923e7 Fix ref latent
oops
2025-12-21 22:53:25 +02:00
kijai dca3106f10 Expose more options, make vid2vid easier 2025-12-20 18:46:32 +02:00
kijai 175418b8d2 Create LongCatAvatar_testing_wip.json 2025-12-20 03:15:24 +02:00
kijai 4a6e2d3c6c Init 2025-12-20 03:14:49 +02:00
11 changed files with 6238 additions and 144 deletions
File diff suppressed because one or more lines are too long
+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.to(self.q_norm.weight.dtype)).to(q.dtype)
# 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.to(self.k_norm.weight.dtype)).to(encoder_k.dtype)
# 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
+119
View File
@@ -0,0 +1,119 @@
import torch
from ..utils import log
import comfy.model_management as mm
from comfy_api.latest import io
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="WanVideoLongCatAvatarExtendEmbeds",
category="WanVideoWrapper",
inputs=[
io.Latent.Input("prev_latents", tooltip="Full previous latents to be used to continue generation, continuation frames are selected based on 'overlap' parameter"),
io.Custom("MULTITALK_EMBEDS").Input("audio_embeds", tooltip="Full length audio embeddings"),
io.Int.Input("num_frames", default=93, min=1, max=256, step=1, tooltip="Number of new frames to generate"),
io.Int.Input("overlap", default=13, min=0, max=16, step=1, tooltip="Number of overlapping frames from previous latents for video continuation, set to 0 for T2V"),
io.Int.Input("frames_processed", default=0, min=0, max=10000, step=1, tooltip="Number of frames already processed in the video, used to select audio features"),
io.Combo.Input("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"),
io.Int.Input("ref_frame_index", default=10, min=0, max=1000, step=1, tooltip="Values between 0 - 24 ensures better consistency, while selecting other ranges (e.g., -10 or 30) helps reduce repeated actions"),
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
],
outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
io.Latent.Output(display_name="samples_slice", tooltip="Sliced latent samples for the new frames"),
],
)
@classmethod
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None) -> io.NodeOutput:
new_audio_embed = audio_embeds.copy()
audio_features = torch.stack(new_audio_embed["audio_features"])
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, :]
prev_samples = prev_latents["samples"].clone()
if overlap != 0:
latent_overlap = (overlap - 1) // 4 + 1
prev_samples = prev_samples[:, :, -latent_overlap:]
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])
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}")
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
longcat_avatar_options = {
"longcat_ref_latent": ref_sample,
"ref_frame_index": ref_frame_index,
"ref_mask_frame_range": ref_mask_frame_range,
}
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
"extra_latents": [{"samples": prev_samples, "index": 0}] if overlap != 0 else None,
"multitalk_embeds": new_audio_embed,
"longcat_avatar_options": longcat_avatar_options,
}
samples_slice = None
if samples is not None:
latent_start_index = (frames_processed - 1) // 4 + 1 if frames_processed > 0 else 0
latent_end_index = latent_start_index + new_latent_frames
samples_slice = samples.copy()
samples_slice["samples"] = samples["samples"][:, :, latent_start_index:latent_end_index].clone()
return io.NodeOutput(embeds, samples_slice)
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)
+35 -2
View File
@@ -1128,7 +1128,7 @@ class WanVideoModelLoader:
scale_weights = {}
if "fp8" in quantization:
for k, v in sd.items():
if k.endswith(".scale_weight"):
if k.endswith(".scale_weight") or k.endswith(".weight_scale"):
is_scaled_fp8 = True
break
@@ -1153,7 +1153,7 @@ class WanVideoModelLoader:
# currently this can be VACE, MTV-Crafter, Lynx or Ovi-audio weights
if extra_model is not None:
for _model in extra_model:
print("Loading extra model: ", _model["path"])
log.info(f"Loading extra model: {_model['path']}")
if gguf:
if not _model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the extra model must also be GGUF quantized, if the main model already has VACE included, you can disconnect the extra module loader")
@@ -1479,6 +1479,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
sd = {k.replace(".weight_scale", ".scale_weight"): v for k, v in sd.items()}
+106 -99
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)
@@ -820,7 +824,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:
@@ -830,7 +834,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]
@@ -841,9 +845,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}")
@@ -869,6 +873,25 @@ class WanVideoSampler:
latent = noise
# LongCat-Avatar
longcat_ref_latent = None
longcat_num_ref_latents = longcat_num_cond_latents = 0
longcat_avatar_options = image_embeds.get("longcat_avatar_options", None)
if longcat_avatar_options is not None:
longcat_ref_latent = longcat_avatar_options.get("longcat_ref_latent", None)
if longcat_ref_latent is not None:
log.info(f"LongCat-Avatar reference latent shape: {longcat_ref_latent.shape}")
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]
longcat_num_ref_latents = longcat_ref_latent.shape[1]
latent_video_length += insert_len
longcat_num_cond_latents = len(clean_latent_indices)
log.info(f"LongCat num_cond_latents: {longcat_num_cond_latents} num_ref_latents: {longcat_num_ref_latents}")
audio_stride = 2 if transformer.is_longcat else 1
#controlnet
controlnet_latents = controlnet = None
if transformer_options is not None:
@@ -1385,27 +1408,30 @@ 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:
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)
@@ -1513,7 +1539,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
@@ -1543,7 +1569,9 @@ 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": longcat_num_cond_latents,
"longcat_num_ref_latents": longcat_num_ref_latents,
"longcat_avatar_options": longcat_avatar_options, # LongCat avatar attention options
"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,
@@ -1576,12 +1604,16 @@ class WanVideoSampler:
if use_fresca:
noise_pred_cond = fourier_filter(noise_pred_cond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
if fantasy_portrait_input is not None and not math.isclose(portrait_cfg[idx], 1.0):
print("Applying Fantasy Portrait CFG...")
base_params["fantasy_portrait_input"] = None
noise_pred_no_portrait, noise_pred_ovi, 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)
noise_pred_no_portrait = noise_pred_no_portrait[0]
return noise_pred_no_portrait + portrait_cfg[idx] * (noise_pred_cond - noise_pred_no_portrait), noise_pred_ovi, [cache_state_cond, cache_state_uncond]
return noise_pred_no_portrait[0] + portrait_cfg[idx] * (noise_pred_cond - noise_pred_no_portrait[0]), noise_pred_ovi, [cache_state_cond, cache_state_uncond]
elif 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]
@@ -1597,12 +1629,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
@@ -1617,8 +1649,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:
@@ -1634,8 +1666,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]
@@ -1647,32 +1679,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):
@@ -1683,7 +1711,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
@@ -1697,7 +1725,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]
@@ -1727,23 +1755,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:
@@ -1786,7 +1814,6 @@ class WanVideoSampler:
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
#torch.cuda.memory._record_memory_history(max_entries=100000)
except:
pass
@@ -1842,7 +1869,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)
@@ -1873,15 +1900,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:
@@ -1895,7 +1922,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
@@ -2265,8 +2292,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
@@ -2308,7 +2335,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)
@@ -3134,28 +3161,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]
@@ -3168,11 +3177,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:
@@ -3242,6 +3247,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:
@@ -3260,8 +3267,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
@@ -3380,6 +3385,7 @@ class WanVideoScheduler:
},
"optional": {
"sigmas": ("SIGMAS", ),
"enhance_hf": ("BOOLEAN", {"default": False, "tooltip": "Enhanced high-frequency denoising schedule"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
@@ -3392,9 +3398,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,
@@ -3479,6 +3485,7 @@ class WanVideoSchedulerv2(WanVideoScheduler):
},
"optional": {
"sigmas": ("SIGMAS", ),
"enhance_hf": ("BOOLEAN", {"default": False, "tooltip": "Enhanced high-frequency denoising schedule"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
+125 -38
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_avatar_options=None, #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
):
@@ -1015,6 +1015,11 @@ class WanAttentionBlock(nn.Module):
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
input_dtype = x.dtype
B, N, C = x.shape
T = num_latent_frames
is_longcat = C == 4096
zero_timestep = len(e) == 2
if zero_timestep: #s2v zero timestep
self.seg_idx = e[1]
@@ -1030,11 +1035,10 @@ 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)
if multitalk_audio_embedding is not None and is_longcat:
audio_shift_mca, audio_scale_mca, audio_gate_mca = self.audio_modulation(e[:, longcat_num_cond_latents:]).unsqueeze(2).chunk(3, dim=-1)
del e
input_dtype = x.dtype
B, N, C = x.shape
T = num_latent_frames
is_longcat = C == 4096
if is_longcat:
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C)
elif use_token_replace:
@@ -1181,22 +1185,67 @@ 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 noise tokens
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# 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)
# 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)
if not longcat_num_cond_latents == num_latent_frames:
# 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 = longcat_avatar_options["ref_mask_frame_range"]
ref_img_index = longcat_avatar_options["ref_frame_index"]
num_ref_latents = 1
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)
# 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)
# 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)
del q, k, v,
del q, k, v
# FETA
if enhance_enabled:
@@ -1263,12 +1312,21 @@ 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)
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 +1340,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:
@@ -1308,7 +1366,7 @@ class WanAttentionBlock(nn.Module):
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=2 if is_longcat else 1)
del mod_x
# gate_mlp
@@ -2128,7 +2186,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 +2203,18 @@ 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+freq_offset + (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)
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+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, 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, 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, longcat_avatar_options=None, # 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
@@ -2559,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
@@ -2567,16 +2638,17 @@ class WanModel(torch.nn.Module):
self.cached_key == cache_key):
freqs = self.cached_freqs
else:
log.info("Generating new RoPE frequencies")
freqs = self.rope_encode_comfy(
F, H, W,
freq_offset=freq_offset,
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
)
log.info("Generated new RoPE frequencies")
if s2v_ref_latent is not None:
freqs_ref = self.rope_encode_comfy(
@@ -2641,8 +2713,8 @@ class WanModel(torch.nn.Module):
if len(t.shape) == 1:
t = t.unsqueeze(1).expand(-1, F) # [B, T]
self.time_embedding.to(torch.float32)
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32).reshape(1, F, -1)
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32)#.reshape(1, F, -1)
e = e0 = e0.reshape(1, F, -1)
if self.audio_model is not None:
#if t.dim() == 1:
@@ -2777,10 +2849,24 @@ 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
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
@@ -2974,7 +3060,8 @@ 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,
longcat_avatar_options=longcat_avatar_options,
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)