Author SHA1 Message Date
kijai fdb23dec7d Update model.py 2026-01-05 22:11:04 +02:00
kijai 07d7d8ca8e remove prints 2026-01-05 22:10:02 +02:00
kijai 01869d4bf5 Merge branch 'main' into longvie2 2026-01-05 18:47:48 +02:00
kijai 55c672028b Merge branch 'main' into longvie2 2025-12-29 15:39:43 +02:00
kijai b551ec9e31 Merge branch 'main' into longvie2 2025-12-29 15:03:53 +02:00
kijai 9f019d7dfb Merge branch 'main' into longvie2 2025-12-23 23:40:25 +02:00
kijai fc5322fae4 Merge branch 'main' into longvie2 2025-12-23 22:04:15 +02:00
kijai 222fc70eb7 Update nodes.py 2025-12-23 17:18:55 +02:00
kijai 8509236da1 init 2025-12-23 14:20:18 +02:00
10 changed files with 902 additions and 5048 deletions
File diff suppressed because one or more lines are too long
+35 -40
View File
@@ -42,43 +42,41 @@ def rotate_half(x):
x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, "... d r -> ... (d r)")
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, split_num=4):
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None):
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
scale = 1.0 / visual_q.shape[-1] ** 0.5
visual_q = visual_q.transpose(1, 2) * scale
visual_q = visual_q * scale
visual_q = visual_q.transpose(1, 2)
ref_k = ref_k.transpose(1, 2)
attn = visual_q @ ref_k.transpose(-2, -1)
if attn_bias is not None:
attn = attn + attn_bias
x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens
B, H, x_seqlens, K = visual_q.shape
x_ref_attn_maps = []
ref_target_masks = ref_target_masks.to(visual_q.dtype)
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
for class_idx, ref_target_mask in enumerate(ref_target_masks):
ref_target_mask = ref_target_mask.view(1, 1, 1, -1)
ref_target_mask = ref_target_mask[None, None, None, ...]
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
x_ref_attnmap = torch.zeros(B, H, x_seqlens, device=visual_q.device, dtype=visual_q.dtype)
chunk_size = min(max(x_seqlens // split_num, 1), x_seqlens)
if mode == 'mean':
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens
elif mode == 'max':
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
for i in range(0, x_seqlens, chunk_size):
end_i = min(i + chunk_size, x_seqlens)
attn_chunk = visual_q[:, :, i:end_i] @ ref_k.permute(0, 2, 3, 1) # B, H, chunk, ref_seqlens
# Apply softmax
attn_max = attn_chunk.max(dim=-1, keepdim=True).values
attn_chunk = (attn_chunk - attn_max).exp()
attn_sum = attn_chunk.sum(dim=-1, keepdim=True)
attn_chunk = attn_chunk / (attn_sum + 1e-8)
# Apply mask and sum
masked_attn = attn_chunk * ref_target_mask
x_ref_attnmap[:, :, i:end_i] = masked_attn.sum(-1) / (ref_target_mask.sum() + 1e-8)
del attn_chunk, masked_attn
# Average across heads
x_ref_attnmap = x_ref_attnmap.mean(dim=1) # B, x_seqlens
x_ref_attn_maps.append(x_ref_attnmap)
del visual_q, ref_k
del attn, x_ref_attn_map_source
return torch.cat(x_ref_attn_maps, dim=0)
return torch.concat(x_ref_attn_maps, dim=0)
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
"""Args:
@@ -131,30 +129,27 @@ class RotaryPositionalEmbedding1D(nn.Module):
query with the same shape as input.
"""
freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
in_dtype = x.dtype
x = x.float()
x_ = x.float()
freqs_cis = freqs_cis.float().to(x.device)
cos = rearrange(freqs_cis.cos(), 'n d -> 1 1 n d')
sin = rearrange(freqs_cis.sin(), 'n d -> 1 1 n d')
cos, sin = freqs_cis.cos(), freqs_cis.sin()
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d')
x_ = (x_ * cos) + (rotate_half(x_) * sin)
# In-place rotation to save memory
x_rotated = rotate_half(x)
x.mul_(cos).add_(x_rotated * sin)
return x.to(in_dtype)
return x_.type_as(x)
class AudioProjModel(nn.Module):
def __init__(
self,
seq_len=5,
seq_len_vf=8,
seq_len_vf=12,
blocks=12,
channels=768,
intermediate_dim=512,
output_dim=768,
context_tokens=32,
norm_output_audio=True,
norm_output_audio=False,
):
super().__init__()
@@ -278,9 +273,9 @@ class SingleStreamMultiAttention(SingleStreamAttention):
def __init__(
self,
dim: int,
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool = True,
encoder_hidden_states_dim: int = 768,
qkv_bias: bool,
class_range: int = 24,
class_interval: int = 4,
attention_mode: str = 'sdpa',
+20 -106
View File
@@ -6,7 +6,7 @@ import numpy as np
from ..latent_preview import prepare_callback
from ..wanvideo.schedulers import get_scheduler
from .multitalk import timestep_transform, add_noise
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap, match_and_blend_colors
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap
from comfy.utils import load_torch_file
from ..nodes_model_loading import load_weights
from ..HuMo.nodes import get_audio_emb_window
@@ -48,13 +48,7 @@ def multitalk_loop(self, **kwargs):
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
elif mode == "skyreelsv3":
num_pseudo_frames = 5
pseudo_frames = reference_keyframes = None
keyframe_index = 0
reference_video = image_embeds.get("reference_video", None)
log.info(f"Multitalk mode: {mode}")
drop_frames = image_embeds.get("drop_frames", 0)
cond_frame = None
offload = image_embeds.get("force_offload", False)
offloaded = False
@@ -68,9 +62,7 @@ def multitalk_loop(self, **kwargs):
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = image_embeds.get("multitalk_start_image", None)
cond_image = original_images.clone() if original_images is not None else None
original_color_reference = cond_image.clone() if cond_image is not None else None
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
@@ -102,6 +94,7 @@ def multitalk_loop(self, **kwargs):
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
cond_frame = None
uni3c_data = None
if uni3c_embeds is not None:
@@ -117,56 +110,9 @@ def multitalk_loop(self, **kwargs):
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame - drop_frames) + 1
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
# If reference_video is provided, extract keyframes from it
if mode == "skyreelsv3" and reference_video is not None:
ref_video_length = reference_video.shape[1] # (C, T, H, W)
if colormatch == "reinhard_torch":
reference_video = match_and_blend_colors(reference_video, original_color_reference, 1.0)
if ref_video_length >= total_frames:
# Reference is long enough - extract keyframes at the expected positions
segment_interval = frame_num - motion_frame - drop_frames
generate_idx = []
current_idx = frame_num - 1
while current_idx < total_frames:
generate_idx.append(min(current_idx, ref_video_length - 1))
current_idx += segment_interval
else:
# Calculate target indices then map to reference video
audio_length = total_frames
generate_idx_target = [0]
segment_interval = frame_num - motion_frame - drop_frames
current_idx = frame_num - 1
while current_idx < audio_length - 1:
generate_idx_target.append(current_idx)
current_idx += segment_interval
if generate_idx_target[-1] != audio_length - 1:
generate_idx_target.append(audio_length - 1)
# Map target indices to reference video
generate_idx_target = np.array(generate_idx_target, dtype=np.int16)
original_max = generate_idx_target[-1]
original_min = generate_idx_target[0]
if original_max > original_min:
generate_idx_float = (generate_idx_target.astype(np.float64) - original_min) * (ref_video_length - 1) / (original_max - original_min)
generate_idx = np.clip(np.round(generate_idx_float), 0, ref_video_length - 1).astype(np.int32).tolist()
else:
generate_idx = [0]
generate_idx = generate_idx[1:]
log.info(f"Reference video ({ref_video_length} frames) mapped to target ({total_frames} frames). Keyframe indices: {generate_idx}")
# Extract keyframes from reference video
# reference_video shape: (C, T, H, W) from nodes.py processing
# Select keyframes and add batch dimension: (C, num_keyframes, H, W) -> (1, C, num_keyframes, H, W)
selected_keyframes = reference_video[:, generate_idx] # (C, num_keyframes, H, W)
reference_keyframes = selected_keyframes.unsqueeze(0).cpu() # (1, C, num_keyframes, H, W)
log.info(f"Extracted {len(generate_idx)} keyframes from provided reference video at indices {generate_idx}, shape: {reference_keyframes.shape}")
log.info(f"Reference video total frames: {reference_video.shape[1]}, will generate {total_frames} total frames with {estimated_iterations} windows")
if frame_num >= total_frames:
arrive_last_frame = True
estimated_iterations = 1
@@ -176,14 +122,6 @@ def multitalk_loop(self, **kwargs):
while True: # start video generation iteratively
self.cache_state = [None, None]
if mode == "skyreelsv3" and reference_keyframes is not None:
clamped_index = min(keyframe_index, reference_keyframes.shape[2] - 1) # Clamp keyframe_index to reuse last keyframe if we run out
pseudo_frames = reference_keyframes[:, :, clamped_index:clamped_index+1].repeat(1, 1, num_pseudo_frames, 1, 1) # Use one keyframe and repeat it 5 times
log.info(f"Window {iteration_count}: using keyframe {clamped_index}/{reference_keyframes.shape[2]-1} for pseudo frames.")
keyframe_index += 1
else:
pseudo_frames = None
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
@@ -195,13 +133,15 @@ def multitalk_loop(self, **kwargs):
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)
audio_embs = torch.cat(audio_embs, dim=0).to(dtype)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(16, latent_frame_num, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
noise = torch.randn(
16, latent_frame_num,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
@@ -258,39 +198,22 @@ def multitalk_loop(self, **kwargs):
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2]
# Prepare pseudo frames if enabled and available from reference_video
if mode == "skyreelsv3" and pseudo_frames is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num-num_pseudo_frames, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames, pseudo_frames.to(device, vae.dtype)], dim=2)
else:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames], dim=2)
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "infinitetalk":
if mode == "multitalk":
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
else:
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
if mode == "skyreelsv3" and pseudo_frames is not None:
# create mask in pixel space, then transform
msk_pixel = torch.ones(1, frame_num, lat_h, lat_w, device=device)
msk_pixel[:, cur_motion_frames_num : -num_pseudo_frames] = 0
msk_pixel = torch.cat([
torch.repeat_interleave(msk_pixel[:, 0:1], repeats=4, dim=1),
msk_pixel[:, 1:],
], dim=1)
msk_pixel = msk_pixel.view(1, msk_pixel.shape[1] // 4, 4, lat_h, lat_w)
msk = msk_pixel.transpose(1, 2).squeeze(0).to(dtype) # 4 T H W
else:
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W
@@ -335,12 +258,11 @@ def multitalk_loop(self, **kwargs):
latent = noise
# injecting motion frames
if not is_first_clip and mode != "infinitetalk":
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent
del motion_add_noise, add_latent
if offloaded:
# Load weights
@@ -448,13 +370,12 @@ def multitalk_loop(self, **kwargs):
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode != "infinitetalk":
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
del motion_add_noise, add_latent
elif mode == "infinitetalk":
else:
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
@@ -464,31 +385,24 @@ def multitalk_loop(self, **kwargs):
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# crop drop_frames from end if enabled
if mode == "skyreelsv3" and drop_frames > 0 and not arrive_last_frame:
videos = videos[:, :-drop_frames]
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
if colormatch == "reinhard_torch":
videos = match_and_blend_colors(videos, original_color_reference, 1.0)
else:
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "infinitetalk":
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
@@ -527,7 +441,7 @@ def multitalk_loop(self, **kwargs):
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count - drop_frames)
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
-88
View File
@@ -462,99 +462,12 @@ class WanVideoImageToVideoMultiTalk:
return (image_embeds, output_path)
class WanVideoImageToVideoSkyreelsv3_audio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the generation"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the generation"}),
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "The number of frames to process at once, should be a value the model is generally good at."}),
"motion_frame": ("INT", {"default": 5, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
"drop_frames": ("INT", {"default": 12, "min": 0, "max": 10000, "step": 1, "tooltip": "Additional frames to drop when advancing the audio window. Higher values = less overlap = faster generation but potentially less smooth transitions."}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"force_offload": ("BOOLEAN", {"default": False, "tooltip": "Whether to force offload the model within the loop for VAE operations, enable if you encounter memory issues."}),
"colormatch": (
[
'disabled',
'reinhard_torch',
'mkl',
'hm',
'reinhard',
'mvgd',
'hm-mvgd-hm',
'hm-mkl-hm',
], {
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
},),
},
"optional": {
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
"reference_video": ("IMAGE", {"tooltip": "Optional: Pre-generated reference video to use for keyframes instead of extracting from first generation. Should be color-matched to source image."}),
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
"output_path": ("STRING", {"default": "", "tooltip": "If set, will save each window's resulting frames to this folder, also DISABLES returning the final video tensor to save memory"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "STRING",)
RETURN_NAMES = ("image_embeds", "output_path")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
def process(self, vae, width, height, frame_window_size, motion_frame, drop_frames, force_offload, colormatch, start_image=None,
tiled_vae=False, clip_embeds=None, mode="multitalk", output_path="", reference_video=None):
H, W = height, width
num_frames = ((frame_window_size - 1) // 4) * 4 + 1
# Resize and rearrange the input image dimensions
if start_image is not None:
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
resized_start_image = resized_start_image * 2 - 1
resized_start_image = resized_start_image.unsqueeze(0)
target_shape = (16, (num_frames - 1) // 4 + 1, height // 8, width // 8)
if output_path:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
output_path = os.path.join(output_path, f"{timestamp}_{mode}_output")
os.makedirs(output_path, exist_ok=True)
processed_reference_video = None
if reference_video is not None:
processed_reference_video = common_upscale(reference_video.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
processed_reference_video = processed_reference_video * 2 - 1
image_embeds = {
"multitalk_sampling": True,
"multitalk_start_image": resized_start_image if start_image is not None else None,
"frame_window_size": num_frames,
"motion_frame": motion_frame,
"drop_frames": drop_frames,
"use_pseudo_frames": True,
"reference_video": processed_reference_video,
"target_h": H,
"target_w": W,
"tiled_vae": tiled_vae,
"force_offload": force_offload,
"vae": vae,
"target_shape": target_shape,
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"colormatch": colormatch,
"multitalk_mode": "skyreelsv3",
"output_path": output_path
}
return (image_embeds, output_path)
NODE_CLASS_MAPPINGS = {
"MultiTalkModelLoader": MultiTalkModelLoader,
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
"Wav2VecModelLoader": Wav2VecModelLoader,
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
"WanVideoImageToVideoSkyreelsv3_audio": WanVideoImageToVideoSkyreelsv3_audio,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -563,5 +476,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
"Wav2VecModelLoader": "Wav2vec2 Model Loader",
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
"WanVideoImageToVideoSkyreelsv3_audio": "WanVideo Long SkyReelsV3 A2V",
}
+4 -92
View File
@@ -2,13 +2,11 @@ import os, gc, math
import torch
import torch.nn.functional as F
import hashlib
from tqdm import tqdm
from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device)
from .taehv import TAEHV
from comfy import model_management as mm
from comfy_api.latest import io
from comfy.utils import ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
import folder_paths
@@ -368,18 +366,10 @@ class WanVideoTextEncode:
cast_dtype = encoder.dtype
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
if hasattr(encoder, 'state_dict'):
model_state_dict = encoder.state_dict
else:
model_state_dict = encoder.model.state_dict()
params_list = list(encoder.model.named_parameters())
pbar = tqdm(params_list, desc="Loading T5 parameters", leave=True)
for name, param in pbar:
for name, param in encoder.model.named_parameters():
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
value = model_state_dict[name]
value = encoder.state_dict[name] if hasattr(encoder, 'state_dict') else encoder.model.state_dict()[name]
set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value)
del model_state_dict
if hasattr(encoder, 'state_dict'):
del encoder.state_dict
mm.soft_empty_cache()
@@ -560,9 +550,6 @@ class WanVideoApplyNAG:
"nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}),
"nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
"inplace": ("BOOLEAN", {"default": True, "tooltip": "If true, modifies tensors in place to save memory. Leads to different numerical results which may change the output slightly."}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
@@ -571,7 +558,7 @@ class WanVideoApplyNAG:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'"
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha, inplace=True):
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha):
prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({
"nag_prompt_embeds": nag_text_embeds["prompt_embeds"],
@@ -579,7 +566,6 @@ class WanVideoApplyNAG:
"nag_scale": nag_scale,
"nag_tau": nag_tau,
"nag_alpha": nag_alpha,
"inplace": inplace,
}
})
return (prompt_embeds_dict_copy,)
@@ -909,8 +895,6 @@ class WanVideoAddStoryMemLatents:
"vae": ("WANVAE",),
"embeds": ("WANVIDIMAGE_EMBEDS",),
"memory_images": ("IMAGE",),
"rope_negative_offset": ("BOOLEAN", {"default": False, "tooltip": "Use positive RoPE frequency offset for the memory latents"}),
"rope_negative_offset_frames": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "RoPE frequency offset for the memory latents"}),
}
}
@@ -919,11 +903,10 @@ class WanVideoAddStoryMemLatents:
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, vae, embeds, memory_images, rope_negative_offset, rope_negative_offset_frames):
def add(self, vae, embeds, memory_images):
updated = dict(embeds)
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
updated["rope_negative_offset_frames"] = rope_negative_offset_frames if rope_negative_offset else 0
return (updated,)
@@ -2070,76 +2053,6 @@ class WanVideoAddTTMLatents:
return (updated,)
#region self-refine-video
class WanVideoSelfRefineVideo(io.ComfyNode):
@classmethod
def define_schema(cls):
# Default values for each range
default_ranges = [
(2, 5, 3), # Range 1
(6, 14, 1), # Range 2
(6, 14, 1), # Range 3
(6, 14, 1), # Range 4
(6, 14, 1), # Range 5
]
options = []
for num_ranges in range(1, 6): # 1 to 5 ranges
range_inputs = []
for i in range(1, num_ranges + 1):
start_default, end_default, steps_default = default_ranges[i - 1]
range_inputs.extend([
io.Int.Input(f"start_step{i}", default=start_default, min=0, max=999, step=1, tooltip=f"Start step for range {i}"),
io.Int.Input(f"end_step{i}", default=end_default, min=0, max=999, step=1, tooltip=f"End step for range {i}"),
io.Int.Input(f"steps_{i}", default=steps_default, min=1, max=100, step=1, tooltip=f"Number of P&P steps for range {i}"),
])
options.append(io.DynamicCombo.Option(
key=str(num_ranges),
inputs=range_inputs
))
return io.Schema(
node_id="WanVideoSelfRefineVideo",
category="WanVideoWrapper",
description="https://github.com/agwmon/self-refine-video - Configure stochastic plan for Perturb-and-Project sampling",
inputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Input("embeds", tooltip="Image embeddings to update"),
io.Float.Input(
"uncertainty_threshold",
default=0.25, min=0.0, max=1.0, step=0.01,
tooltip="Lower values make it harder for regions to be considered \"certain\", meaning more pixels will continue being refined. Higher values make it easier to lock in pixels early."
),
io.Float.Input("certain_percentage", default=0.999, min=0.0, max=1.0, step=0.001, tooltip="Higher values = stricter requirement = fewer early stops = more iterations"),
io.DynamicCombo.Input("num_ranges", options=options, display_name="Number of Ranges", tooltip="Number of step ranges to configure for the stochastic plan"),
],
outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Updated image embeddings with self-refine parameters"),
],
)
@classmethod
def execute(cls, embeds, uncertainty_threshold, certain_percentage, num_ranges) -> io.NodeOutput:
updated = dict(embeds)
updated["self_refine_uncertainty_threshold"] = uncertainty_threshold
updated["self_refine_certain_percentage"] = certain_percentage
# Build stochastic plan from the dynamic inputs in list format: [(start, end, steps), ...]
stochastic_plan = []
range_keys = sorted([k for k in num_ranges.keys() if k.startswith('start_step')])
for start_key in range_keys:
i = start_key.replace('start_step', '')
start = num_ranges.get(f"start_step{i}")
end = num_ranges.get(f"end_step{i}")
steps = num_ranges.get(f"steps_{i}")
if start is not None and end is not None and steps is not None:
stochastic_plan.append((start, end, steps))
updated["stochastic_plan"] = stochastic_plan
return io.NodeOutput(updated)
#region VideoDecode
class WanVideoDecode:
@classmethod
@@ -2394,7 +2307,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
"WanVideoSelfRefineVideo": WanVideoSelfRefineVideo,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+70 -62
View File
@@ -52,22 +52,9 @@ def update_folder_names_and_paths(key, targets=[]):
log.warning(f"Unknown file list already present on key {key}: {base}")
update_folder_names_and_paths("unet_gguf", ["diffusion_models", "unet"])
try:
from comfy.latent_formats import Wan21, Wan22
latent_format = Wan21
except: #for backwards compatibility
log.warning("WARNING: Wan21 latent format not found, update ComfyUI for better live video preview")
from comfy.latent_formats import HunyuanVideo
latent_format = HunyuanVideo
class WanVideoModel(torch.nn.Module):
def __init__(self, model_config, transformer, device=None):
super().__init__()
self.latent_format = model_config.latent_format
self.model_config = model_config
self.device = device
self.current_patcher = None
self.diffusion_model = transformer
class WanVideoModel(comfy.model_base.BaseModel):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.pipeline = {}
def __getitem__(self, k):
@@ -76,11 +63,24 @@ class WanVideoModel(torch.nn.Module):
def __setitem__(self, k, v):
self.pipeline[k] = v
try:
from comfy.latent_formats import Wan21, Wan22
latent_format = Wan21
except: #for backwards compatibility
log.warning("WARNING: Wan21 latent format not found, update ComfyUI for better live video preview")
from comfy.latent_formats import HunyuanVideo
latent_format = HunyuanVideo
class WanVideoModelConfig:
def __init__(self, latent_format=latent_format):
def __init__(self, dtype, latent_format=latent_format):
self.unet_config = {}
self.unet_extra_config = {}
self.latent_format = latent_format
#self.latent_format.latent_channels = 16
self.manual_cast_dtype = dtype
self.sampling_settings = {"multiplier": 1.0}
self.memory_usage_factor = 2.0
self.unet_config["disable_unet_model_creation"] = True
def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]):
filtered_dict = {}
@@ -810,6 +810,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
cnt = 0
block_idx = vace_block_idx = None
if gguf:
@@ -829,7 +830,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
all_tensors.extend(r.tensors)
for tensor in all_tensors:
name = rename_fuser_block(tensor.name)
if "glob" not in name and "multitalk_audio_proj" not in name and "audio_proj" in name:
if "glob" not in name and "audio_proj" in name:
name = name.replace("audio_proj", "multitalk_audio_proj")
load_device = device
if "vace_blocks." in name:
@@ -919,7 +920,9 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
load_device = offload_device
# Set tensor to device
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=value)
pbar.update(1)
cnt += 1
if cnt % 100 == 0:
pbar.update(100)
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
memory_on_device = get_module_memory_mb_per_device(transformer)
@@ -928,8 +931,6 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
for dev, mem_mb in memory_on_device.items():
log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB")
if hasattr(pbar, "_last_sent_value"):
pbar._last_sent_value = -1
pbar.update_absolute(0)
def patch_control_lora(transformer, device):
@@ -1511,45 +1512,7 @@ class WanVideoModelLoader:
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
# LongCat Avatar
if "multitalk_audio_proj.proj1.weight" in sd and "blocks.0.audio_cross_attn.q_norm.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
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()
transformer.multitalk_audio_proj = multitalk_proj_model
# SkyreelsV3
elif "blocks.1.audio_cross_attn.kv_linear.weight" in sd and "audio_proj.proj1.weight" in sd:
sd = {k.replace("audio_proj", "multitalk_audio_proj"): v for k, v in sd.items()}
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention, AudioProjModel
from .wanvideo.modules.model import WanLayerNorm
for block in transformer.blocks:
with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
transformer.multitalk_audio_proj = AudioProjModel()
elif multitalk_model is not None:
if multitalk_model is not None:
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
log.info(f"{multitalk_model_type} detected, patching model...")
@@ -1566,7 +1529,15 @@ class WanVideoModelLoader:
for block in transformer.blocks:
with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
block.audio_cross_attn = SingleStreamMultiAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
transformer.multitalk_audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type
@@ -1584,6 +1555,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()}
@@ -1609,7 +1613,11 @@ class WanVideoModelLoader:
transformer.text_projection = nn.Sequential(nn.Linear(sd["text_projection.0.weight"].shape[1], text_dim), nn.GELU(approximate='tanh'), nn.Linear(text_dim, text_dim))
latent_format=Wan22 if dim == 3072 else Wan21
comfy_model = WanVideoModel(WanVideoModelConfig(latent_format=latent_format), device=device, transformer=transformer)
comfy_model = WanVideoModel(
WanVideoModelConfig(base_dtype, latent_format=latent_format),
model_type=comfy.model_base.ModelType.FLOW,
device=device,
)
# SteadyDancer
if "condition_embedding_align.cross_attn.in_proj_bias" in sd:
+1 -106
View File
@@ -185,7 +185,6 @@ class WanVideoSampler:
is_pusa = "pusa" in sample_scheduler.__class__.__name__.lower()
if scheduler != "multitalk":
scheduler_step_args = {"generator": seed_g}
step_sig = inspect.signature(sample_scheduler.step)
for arg in list(scheduler_step_args.keys()):
@@ -226,7 +225,6 @@ class WanVideoSampler:
#I2V
story_mem_latents = image_embeds.get("story_mem_latents", None)
image_cond = image_embeds.get("image_embeds", None)
image_cond_mask = None
if image_cond is not None:
if transformer.in_dim == 16:
raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models")
@@ -1162,26 +1160,6 @@ class WanVideoSampler:
prev_ones = torch.ones(20, *prev_latents.shape[1:], device=device, dtype=dtype)
dual_control_input["prev_latent"] = torch.cat([prev_ones, prev_latents]).unsqueeze(0)
# Self-refine video
self_refine_video_enabled = False
stochastic_step_map = {}
self_refine_uncertainty_threshold = image_embeds.get("self_refine_uncertainty_threshold", None)
if self_refine_uncertainty_threshold is not None and context_options is None:
self_refine_video_enabled = True
stochastic_plan = image_embeds.get("stochastic_plan", [(2, 5, 3), (6, 14, 1)])
certain_percentage = image_embeds.get("self_refine_certain_percentage", 0.999)
# Build step map: denoising_iteration -> num_anneal_steps
for start, end, steps in stochastic_plan:
for idx in range(start, end + 1):
stochastic_step_map[idx] = steps
log.info("Self-refine video enabled:")
log.info(f" Uncertainty threshold: {self_refine_uncertainty_threshold}")
log.info(f" Certain percentage: {certain_percentage}")
log.info(f" Stochastic plan (steps start-end: P&P iterations):")
for start, end, steps in stochastic_plan:
log.info(f" Steps {start}-{end}: {steps} P&P iteration{'s' if steps != 1 else ''}")
log.info(f" Total denoising steps with P&P: {len(stochastic_step_map)}")
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1510,9 +1488,7 @@ class WanVideoSampler:
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
"scail_input": scail_data_in, # SCAIL input
"dual_control_input": dual_control_in, # LongVie2 dual control input
"transformer_options": transformer_options,
"rope_negative_offset": image_embeds.get("rope_negative_offset_frames", 0), # StoryMem rope negative offset
"num_memory_frames": story_mem_latents.shape[1] if story_mem_latents is not None else 0, # StoryMem memory frames
"transformer_options": transformer_options
}
batch_size = 1
@@ -1818,87 +1794,6 @@ class WanVideoSampler:
latent_flipped = torch.flip(latent, dims=[1])
latent_model_input_flipped = latent_flipped.to(device)
if self_refine_video_enabled:
certain_percentage: float = image_embeds.get("self_refine_certain_percentage", 0.999) # if the certain region is more than this percentage, set certain_flag to True
# Get number of P&P iterations for this noise level
current_num_anneal_steps = stochastic_step_map.get(idx, 0)
in_stoch = current_num_anneal_steps > 0
# m = number of vector-predictions at this step
# If stochastic: 1 (initial) + anneal_steps (re-noised predictions)
m = current_num_anneal_steps + 1 if in_stoch else 1
if sample_scheduler._step_index is None:
tqdm.write("Initializing step index for self-refine scheduler")
sample_scheduler._init_step_index(t)
sigma = sample_scheduler.sigmas[sample_scheduler._step_index].to(device)
sigma_next = sample_scheduler.sigmas[sample_scheduler._step_index + 1].to(device)
buffer = [None] # tuple (certain_mask, pred_original_sample, latents_next)
certain_mask = None
certain_flag = False
m = current_num_anneal_steps + 1 if in_stoch else 1
stoch_bar_context = tqdm(total=current_num_anneal_steps) if in_stoch else nullcontext()
with stoch_bar_context as stoch_bar:
for ii in range(m):
if certain_flag:
latent = buffer[-1][2]
break
if ii == 0: # first prediction step
latent_model_input = latent.to(device)
else: # Perturbation step in Eq. 6 of the paper
noise = torch.randn(latent.shape, generator=seed_g, device=torch.device("cpu"), dtype=latent.dtype)
latent = (1.0 - sigma) * buffer[-1][1] + sigma * noise.to(latent)
latent_model_input = latent.to(device)
latent_model_input_ovi = latent_ovi.to(device) if latent_ovi is not None else None
timestep = torch.tensor([t]).to(device)
latent = latent.to(device)
noise_pred, noise_pred_ovi, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input,
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,
wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data, latent_model_input_ovi=latent_model_input_ovi, flashvsr_LQ_latent=flashvsr_LQ_latent,
)
pred_original_sample = latent - sigma * noise_pred
latent_next = latent + (sigma_next - sigma) * noise_pred
if in_stoch: # stochastic sampling
if buffer[-1] is not None: # if buffer is not empty
# Compute uncertainty per spatial-temporal location
diff = pred_original_sample - buffer[-1][1]
uncertainty = torch.sqrt(torch.sum(diff ** 2, dim=0)) / 16 # f, h, w
certain_mask = uncertainty < self_refine_uncertainty_threshold # certain region, e.g., background
if buffer[-1][0] is not None:
certain_mask = certain_mask | buffer[-1][0] # update certain mask (union)
# if the certain region is more than this percentage, set certain_flag to True
if certain_mask.sum() / certain_mask.numel() > certain_percentage:
certain_flag = True
tqdm.write(f"{ii}/{current_num_anneal_steps}: Certain region is more than {certain_percentage}, set certain_flag to True")
certain_mask_float = certain_mask.float().unsqueeze(0) # 1, f, h, w to broadcast with c, f, h, w
latent_next = certain_mask_float * buffer[-1][2] + (1.0 - certain_mask_float) * latent_next
pred_original_sample = certain_mask_float * buffer[-1][1] + (1.0 - certain_mask_float) * pred_original_sample
buffer.append([certain_mask, pred_original_sample, latent_next])
if ii == m - 1: # finish the last step in P&P iterations
latent = buffer[-1][2]
else: # base ODE step
latent = latent_next
if in_stoch and ii > 0 and stoch_bar is not None:
stoch_bar.update()
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
sample_scheduler._step_index += 1
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
else:
self.noise_front_pad_num = 0
#InfiniteTalk first frame handling
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.4.7"
version = "1.4.5"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+3 -60
View File
@@ -7,14 +7,9 @@ from pathlib import Path
import gc
import types, collections
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
from comfy.model_patcher import get_key_weight
from comfy.model_patcher import get_key_weight, string_to_seed
from comfy.lora import calculate_weight
try:
from comfy.utils import string_to_seed
except:
from comfy.model_patcher import string_to_seed
from comfy.float import stochastic_rounding
from .custom_linear import remove_lora_from_module
import folder_paths
@@ -195,9 +190,9 @@ def set_module_tensor_to_device(module, tensor_name, device, value=None, dtype=N
device = device_quantization
if is_buffer:
module._buffers[tensor_name] = new_value
elif value is not None or not check_device_same(device, module._parameters[tensor_name].device):
elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device):
param_cls = type(module._parameters[tensor_name])
new_value = param_cls(new_value, requires_grad=False)
new_value = param_cls(new_value, requires_grad=False).to(device)
module._parameters[tensor_name] = new_value
#if device != "cpu":
@@ -723,55 +718,3 @@ def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.
if not t == 1.0:
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
return model_output
def match_and_blend_colors(
source_chunk: torch.Tensor, # (C, T, H, W), range [-1, 1]
reference_image: torch.Tensor, # (C, 1, H, W), range [-1, 1]
strength: float,
) -> torch.Tensor:
import kornia
if strength == 0.0:
return source_chunk
source_chunk = source_chunk.unsqueeze(0) # (1, C, T, H, W)
# shapes
B, C, T, H, W = source_chunk.shape
input_dtype = source_chunk.dtype
# [-1,1] -> [0,1]
src_01 = (source_chunk + 1.0) * 0.5
ref_01 = (reference_image + 1.0) * 0.5
src32 = src_01.to(torch.float32)
ref32 = ref_01.to(torch.float32)
# (B, C, T, H, W) -> (B*T, C, H, W)
src_bt = src32.permute(0, 2, 1, 3, 4).contiguous().view(B * T, C, H, W)
ref_bchw = ref32[:, :, 0, :, :].contiguous()
# RGB->Lab
src_lab = kornia.color.rgb_to_lab(src_bt) # (B*T, C, H, W)
ref_lab = kornia.color.rgb_to_lab(ref_bchw) # (B, C, H, W)
src_lab_flat = src_lab.view(B * T, C, -1) # (B*T, C, HW)
ref_lab_flat = ref_lab.view(B, C, -1) # (B, C, HW)
src_std, src_mean = torch.std_mean(src_lab_flat, dim=-1, keepdim=True, unbiased=False)
ref_std, ref_mean = torch.std_mean(ref_lab_flat, dim=-1, keepdim=True, unbiased=False)
src_std = src_std.clamp_min_(1e-6)
ref_mean_bt = ref_mean.repeat_interleave(T, dim=0) # (B*T, C, 1)
ref_std_bt = ref_std.repeat_interleave(T, dim=0) # (B*T, C, 1)
corrected_lab_flat = (src_lab_flat - src_mean) * (ref_std_bt / src_std) + ref_mean_bt
corrected_lab = corrected_lab_flat.view(B * T, C, H, W)
# Lab->RGB
corrected_rgb_01 = kornia.color.lab_to_rgb(corrected_lab) # (B*T, C, H, W)
blended_rgb_01 = (1.0 - strength) * src_bt + strength * corrected_rgb_01
# (B, C, T, H, W)
blended_rgb_01 = blended_rgb_01.view(B, T, C, H, W).permute(0, 2, 1, 3, 4).contiguous()
# [0,1] -> [-1,1]
return (blended_rgb_01 * 2.0 - 1.0)[0].to(dtype=input_dtype)
+26 -60
View File
@@ -558,53 +558,39 @@ class WanSelfAttention(nn.Module):
# output
return self.o(x.flatten(2))
def nag_attention(self, b, n, d, q, context, nag_context=None):
k_positive = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_positive = self.v(context).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode, heads=self.num_heads)
del k_positive, v_positive
k_negative = self.norm_k(self.k(nag_context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_negative = self.v(nag_context).view(b, -1, n, d)
x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode, heads=self.num_heads)
del k_negative, v_negative
return x_positive.flatten(2), x_negative.flatten(2)
def normalized_attention_guidance(self, x_positive, x_negative,nag_params={}):
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
# NAG text attention
context_positive = context
context_negative = nag_context
nag_scale = nag_params['nag_scale']
nag_alpha = nag_params['nag_alpha']
nag_tau = nag_params['nag_tau']
inplace = nag_params.get('inplace', True)
if inplace:
nag_guidance = x_negative.mul_(nag_scale - 1).neg_().add_(x_positive, alpha=nag_scale)
else:
k_positive = self.norm_k(self.k(context_positive).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_positive = self.v(context_positive).view(b, -1, n, d)
k_negative = self.norm_k(self.k(context_negative).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_negative = self.v(context_negative).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode, heads=self.num_heads)
x_positive = x_positive.flatten(2)
x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode, heads=self.num_heads)
x_negative = x_negative.flatten(2)
nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1)
del x_negative
norm_positive = torch.norm(x_positive, p=1, dim=-1, keepdim=True)
norm_guidance = torch.norm(nag_guidance, p=1, dim=-1, keepdim=True)
scale = norm_guidance / norm_positive
torch.nan_to_num_(scale, nan=10.0)
scale = torch.nan_to_num(scale, nan=10.0)
mask = scale > nag_tau
del scale
adjustment = (norm_positive * nag_tau) / (norm_guidance + 1e-7)
del norm_positive, norm_guidance
nag_guidance.mul_(torch.where(mask, adjustment, 1.0))
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance)
del mask, adjustment
if inplace:
nag_guidance.sub_(x_positive).mul_(nag_alpha).add_(x_positive)
else:
nag_guidance = nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
del x_positive
return nag_guidance
return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
class LoRALinearLayer(nn.Module):
def __init__(
@@ -662,10 +648,7 @@ class WanT2VCrossAttention(WanSelfAttention):
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
if nag_context is not None:
x_positive, x_negative = self.nag_attention(b, n, d, q, context, nag_context)
del q
x = self.normalized_attention_guidance(x_positive, x_negative, nag_params)
del x_positive, x_negative
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
else:
if is_longcat:
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype).view(b, -1, n, d)).to(x.dtype)
@@ -769,8 +752,7 @@ class WanI2VCrossAttention(WanSelfAttention):
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
x_text.add_(img_x)
x = x_text
x = x_text + img_x
else:
x = x_text
@@ -1298,7 +1280,7 @@ class WanAttentionBlock(nn.Module):
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x.addcmul_(y, gate_msa)
x = x.addcmul(y, gate_msa)
del y, gate_msa
# cross-attention & ffn function
@@ -1328,10 +1310,11 @@ class WanAttentionBlock(nn.Module):
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
return x, x_ip, lynx_ref_feature, x_ovi
else:
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,
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, longcat_num_cond_latents=longcat_num_cond_latents).to(input_dtype)
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):
@@ -1345,8 +1328,7 @@ class WanAttentionBlock(nn.Module):
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.add_(x_audio, alpha=audio_scale)
del x_audio
x = x.add(x_audio, alpha=audio_scale)
# MTV-Crafter Motion Attention
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
@@ -2206,7 +2188,7 @@ 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,
ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=5):
ref_frame_index=10, longcat_num_ref_latents=0):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
@@ -2230,15 +2212,6 @@ class WanModel(torch.nn.Module):
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)
elif num_memory_frames > 0 and rope_negative_offset > 0:
# Negative RoPE shift for memory frames
# Memory frames get negative indices: {-f_m*S, -(f_m-1)*S, ..., -S}
# Current video frames start from 0: {0, 1, ..., f-1}
memory_indices = torch.arange(-num_memory_frames * rope_negative_offset, 0, rope_negative_offset, dtype=dtype, device=device)
current_indices = torch.arange(0, steps_t - num_memory_frames, dtype=dtype, device=device)
grid_t = torch.cat([memory_indices, current_indices], dim=0)
log.info(f"{num_memory_frames} memory frames, temporal rope positions: {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+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
@@ -2348,8 +2321,6 @@ class WanModel(torch.nn.Module):
scail_input=None, # SCAIL pose
dual_control_input=None, # LongVie2 dual controlnet
transformer_options={},
rope_negative_offset=0,
num_memory_frames=0,
):
r"""
Forward pass through the diffusion model
@@ -2594,7 +2565,6 @@ class WanModel(torch.nn.Module):
scail_x = [u.flatten(2).transpose(1, 2) * scail_input.get("pose_strength", 1) for u in scail_x]
x = [torch.cat([u, v], dim=1) for u, v in zip(x, scail_x)]
seq_len += scail_x[0].shape[1]
del scail_x
pose_frame_shape = scail_pose_latents.shape
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
@@ -2673,8 +2643,6 @@ class WanModel(torch.nn.Module):
self.rope_embedder.k,
tuple(ntk_alphas),
longcat_num_ref_latents,
rope_negative_offset,
num_memory_frames,
)
# Check cache using key comparison
@@ -2690,8 +2658,6 @@ class WanModel(torch.nn.Module):
ref_frame_shape=ref_frame_shape,
pose_frame_shape=pose_frame_shape,
longcat_num_ref_latents=longcat_num_ref_latents,
rope_negative_offset=rope_negative_offset,
num_memory_frames=num_memory_frames,
device=x.device,
dtype=x.dtype
)