diff --git a/nodes.py b/nodes.py index 485cb33..30817d9 100644 --- a/nodes.py +++ b/nodes.py @@ -5,7 +5,8 @@ import torch import numpy as np from tqdm import tqdm import webrtcvad - +from comfy.utils import common_upscale +import datetime class GetFloatByIndex: @@ -171,8 +172,7 @@ class InfiniteTalkMultiImage(): @staticmethod def calculate_big_loops(**kwargs): fps = Decimal("25") - step = Decimal("0.04") - + auto_start_time_list = kwargs.get("auto_start_time_list", []) print("auto_start_time_list:", auto_start_time_list) @@ -185,7 +185,7 @@ class InfiniteTalkMultiImage(): raw_text = kwargs.get("prompt_list_input", "") text_list = [line.strip() for line in raw_text.split("\n") if line.strip()] - # =============== 图片 + start_time =============== + # =============== 图片 + start_time 数据收集 =============== pairs = [] for i in range(1, 21): @@ -200,26 +200,54 @@ class InfiniteTalkMultiImage(): if not pairs: return [], [], [], text_list - start_times = [p[0] for p in pairs] # Decimal - images = [p[1] for p in pairs] + # 分离出初步的时间和图片列表 + raw_start_times = [p[0] for p in pairs] + raw_images = [p[1] for p in pairs] + # 如果有自动时间列表,覆盖手动输入的时间 if auto_start_time_list is not None and len(auto_start_time_list) > 0: - start_times = [Decimal(str(t)) for t in auto_start_time_list] + raw_start_times = [Decimal(str(t)) for t in auto_start_time_list] - min_len = min(len(images), len(start_times)) + # 确保图片和时间列表长度一致(取最小值) + min_len = min(len(raw_images), len(raw_start_times)) + raw_images = raw_images[:min_len] + raw_start_times = raw_start_times[:min_len] - print("min_len", min_len) + # =============== 核心修改:根据音频时长过滤/对齐 =============== + images = [] + start_times = [] - images = images[:min_len] - start_times = start_times[:min_len] + # 遍历所有待选片段,只保留开始时间早于音频总时长的片段 + for img, st in zip(raw_images, raw_start_times): + if st < audio_duration: + images.append(img) + start_times.append(st) + else: + # 一旦发现起始时间超过或等于音频时长,后面的通常也都不需要了 + # 这里不做 break 是为了防止输入时间乱序的情况,虽然通常是顺序的 + continue + + # 如果过滤后为空(例如音频极短,第一张图开始时间都比音频长) + if not start_times: + # 返回空列表或默认处理,防止后续报错 + # 这里选择返回空,或者您可以选择至少保留第一张图给 0 帧 + return [], [], [], [] # =============== frame count(真实差值) =============== frame_count_list = [] for idx, st in enumerate(start_times): + # 确定当前片段的结束时间 if idx < len(start_times) - 1: + # 如果不是最后一张,结束时间通常是下一张的开始时间 end = start_times[idx + 1] + + # 【保护措施】如果下一张图的开始时间意外超过了音频时长(虽然前面过滤过,但防止乱序) + # 或者逻辑上我们希望片段不超出音频范围 + if end > audio_duration: + end = audio_duration else: + # 如果是最后一张,结束时间就是音频总时长 end = audio_duration dur = end - st @@ -233,6 +261,7 @@ class InfiniteTalkMultiImage(): real_start_time_list = [float(t) for t in start_times] # =============== prompt 对齐 =============== + # 此时 images 已经是被 audio_duration 过滤后的列表 n_images = len(images) n_texts = len(text_list) @@ -240,14 +269,12 @@ class InfiniteTalkMultiImage(): # 如果文本数量不够,用空字符串填充 text_list = text_list + [""] * (n_images - n_texts) elif n_texts > n_images: - # 如果文本数量多于图片数量,截断多余部分 + # 如果文本数量多于(过滤后的)图片数量,截断多余部分 text_list = text_list[:n_images] return images, frame_count_list, real_start_time_list, text_list - - class InfiniteTalkEmbedsSlice: @classmethod def INPUT_TYPES(s): @@ -538,28 +565,69 @@ class AudioSmartSlice: import os +import glob +import subprocess +import shutil +import tempfile +import torch +import torchaudio +from PIL import Image +import numpy as np + +import folder_paths +# ========== 全部必要 imports ========== +import os +import glob import shutil import subprocess import tempfile -import folder_paths -import sys -import glob +import random +import math +import numpy as np +from PIL import Image, ImageDraw, ImageFilter, ImageChops, ImageEnhance + +import torch +import torchaudio + +import folder_paths + +# ========== VideoFromPathsAndAudio 节点(完整) ========== class VideoFromPathsAndAudio: def __init__(self): pass - + @classmethod def INPUT_TYPES(s): return { "required": { "image_paths": ("STRING", {"multiline": True, "default": "", "placeholder": "每行一个路径 (文件夹或文件)"}), "frame_counts": ("STRING", {"multiline": True, "default": "", "placeholder": "每行对应路径的帧数"}), - "audio": ("AUDIO",), - # 新增 fps 参数,默认 25 - "fps": ("INT", {"default": 25, "min": 1, "max": 120, "step": 1, "display": "number"}), + "audio": ("AUDIO",), + "fps": ("INT", {"default": 25, "min": 1, "max": 120, "step": 1}), "filename_prefix": ("STRING", {"default": "video_output"}), "format": (["mp4", "mkv", "mov"],), + + # 新增过渡参数(严格保持你要求的名字和位置) + "enable_transition": ("BOOLEAN", {"default": True}), + "transition_seconds": ("FLOAT", {"default": 0.5, "min": 0.05, "max": 5.0, "step": 0.01}), + "transition_type": ( + [ + "crossfade", + "left_wipe", + "right_wipe", + # "up_wipe", + # "down_wipe", + "slide_left", + "slide_right", + # "slide_up", + # "slide_down", + "zoom_in", + "zoom_out", + "soft_blur", + "random", + ], + ), }, } @@ -568,189 +636,933 @@ class VideoFromPathsAndAudio: FUNCTION = "execute_video_synth" CATEGORY = "Custom/Video" + # ---------------------------- + # 打开两张图片并保持大小一致 + # ---------------------------- + def _open_and_match(self, pathA, pathB): + A = Image.open(pathA).convert("RGB") + B = Image.open(pathB).convert("RGB") + if A.size != B.size: + B = B.resize(A.size, Image.LANCZOS) + return A, B + + # ---------------------------- + # 缓动函数(Easing Functions)- 让过渡更自然 + # ---------------------------- + def _ease_in_out_cubic(self, t): + """平滑的缓入缓出曲线""" + if t < 0.5: + return 4 * t * t * t + else: + return 1 - pow(-2 * t + 2, 3) / 2 + + def _ease_out_cubic(self, t): + """缓出曲线""" + return 1 - pow(1 - t, 3) + + def _ease_in_out_quad(self, t): + """二次缓入缓出""" + if t < 0.5: + return 2 * t * t + else: + return 1 - pow(-2 * t + 2, 2) / 2 + + def _smooth_step(self, t): + """平滑步进函数""" + return t * t * (3 - 2 * t) + + # ---------------------------- + # 过渡:crossfade(优化版 - 使用缓动和更好的混合) + # ---------------------------- + def _crossfade(self, A, B, t): + # 使用缓动函数让过渡更平滑 + t_smooth = self._ease_in_out_cubic(t) + # 使用更自然的混合 + return Image.blend(A, B, t_smooth) + + # ---------------------------- + # 过渡:left/right/up/down wipe(优化版 - 添加渐变边缘) + # ---------------------------- + def _wipe(self, A, B, t, direction): + w, h = A.size + # 使用缓动函数 + t_smooth = self._ease_in_out_quad(t) + + # 创建渐变边缘的 mask(让擦除更柔和) + mask = Image.new("L", (w, h), 0) + draw = ImageDraw.Draw(mask) + gradient_width = max(20, min(w, h) // 10) # 渐变宽度为画面的10%,最少20像素 + + if direction == "left": + cut = int(w * t_smooth) + # 填充已擦除区域(完全显示B的部分) + if cut > gradient_width: + draw.rectangle([0, 0, cut - gradient_width, h], fill=255) + # 创建渐变区域 + for x in range(max(0, cut - gradient_width), min(w, cut + gradient_width)): + if x < cut: + alpha = 255 + else: + # 渐变边缘 + dist = x - cut + alpha = max(0, 255 - int(255 * dist / gradient_width)) + draw.rectangle([x, 0, x + 1, h], fill=alpha) + elif direction == "right": + cut = int(w * t_smooth) + # 填充已擦除区域 + if cut > gradient_width: + draw.rectangle([w - cut + gradient_width, 0, w, h], fill=255) + # 创建渐变区域 + for x in range(max(0, w - cut - gradient_width), min(w, w - cut + gradient_width)): + if x >= w - cut: + alpha = 255 + else: + dist = (w - cut) - x + alpha = max(0, 255 - int(255 * dist / gradient_width)) + draw.rectangle([x, 0, x + 1, h], fill=alpha) + elif direction == "up": + cut = int(h * t_smooth) + # 填充已擦除区域 + if cut > gradient_width: + draw.rectangle([0, 0, w, cut - gradient_width], fill=255) + # 创建渐变区域 + for y in range(max(0, cut - gradient_width), min(h, cut + gradient_width)): + if y < cut: + alpha = 255 + else: + dist = y - cut + alpha = max(0, 255 - int(255 * dist / gradient_width)) + draw.rectangle([0, y, w, y + 1], fill=alpha) + else: # down + cut = int(h * t_smooth) + # 填充已擦除区域 + if cut > gradient_width: + draw.rectangle([0, h - cut + gradient_width, w, h], fill=255) + # 创建渐变区域 + for y in range(max(0, h - cut - gradient_width), min(h, h - cut + gradient_width)): + if y >= h - cut: + alpha = 255 + else: + dist = (h - cut) - y + alpha = max(0, 255 - int(255 * dist / gradient_width)) + draw.rectangle([0, y, w, y + 1], fill=alpha) + + return Image.composite(B, A, mask) + + # ---------------------------- + # 过渡:slide(优化版 - 使用缓动和边缘混合) + # ---------------------------- + def _slide(self, A, B, t, direction): + w, h = A.size + # 使用缓动函数让滑动更平滑 + t_smooth = self._ease_in_out_cubic(t) + + canvas = Image.new("RGB", (w, h)) + blend_width = max(10, min(w, h) // 20) # 混合边缘宽度 + + if direction in ("left", "right"): + offset = int(w * t_smooth) + if direction == "left": + # B从右往左推进,A往左退出 + if offset > 0: + # B的右侧部分 + B_part = B.crop((w - offset, 0, w, h)) + if offset < w: + B_part = B_part.resize((offset, h), Image.LANCZOS) + canvas.paste(B_part, (0, 0)) + + # A的剩余部分 + if offset < w: + A_part = A.crop((0, 0, w - offset, h)) + canvas.paste(A_part, (offset, 0)) + + # 在边界处添加混合效果 + if offset > blend_width and offset < w: + blend_mask = Image.new("L", (blend_width, h)) + blend_draw = ImageDraw.Draw(blend_mask) + for x in range(blend_width): + alpha = int(255 * x / blend_width) + blend_draw.rectangle([x, 0, x + 1, h], fill=alpha) + blend_region = Image.composite( + B.crop((w - offset, 0, w - offset + blend_width, h)), + A.crop((offset - blend_width, 0, offset, h)), + blend_mask + ) + canvas.paste(blend_region, (offset - blend_width, 0)) + else: # right + if offset > 0: + B_part = B.crop((0, 0, offset, h)) + if offset < w: + B_part = B_part.resize((offset, h), Image.LANCZOS) + canvas.paste(B_part, (w - offset, 0)) + + if offset < w: + A_part = A.crop((offset, 0, w, h)) + canvas.paste(A_part, (0, 0)) + + if offset > blend_width and offset < w: + blend_mask = Image.new("L", (blend_width, h)) + blend_draw = ImageDraw.Draw(blend_mask) + for x in range(blend_width): + alpha = int(255 * (blend_width - x) / blend_width) + blend_draw.rectangle([x, 0, x + 1, h], fill=alpha) + blend_region = Image.composite( + A.crop((offset - blend_width, 0, offset, h)), + B.crop((0, 0, blend_width, h)), + blend_mask + ) + canvas.paste(blend_region, (offset - blend_width, 0)) + else: + offset = int(h * t_smooth) + if direction == "up": + if offset > 0: + B_part = B.crop((0, h - offset, w, h)) + if offset < h: + B_part = B_part.resize((w, offset), Image.LANCZOS) + canvas.paste(B_part, (0, 0)) + + if offset < h: + A_part = A.crop((0, 0, w, h - offset)) + canvas.paste(A_part, (0, offset)) + + if offset > blend_width and offset < h: + blend_mask = Image.new("L", (w, blend_width)) + blend_draw = ImageDraw.Draw(blend_mask) + for y in range(blend_width): + alpha = int(255 * y / blend_width) + blend_draw.rectangle([0, y, w, y + 1], fill=alpha) + blend_region = Image.composite( + B.crop((0, h - offset, w, h - offset + blend_width)), + A.crop((0, offset - blend_width, w, offset)), + blend_mask + ) + canvas.paste(blend_region, (0, offset - blend_width)) + else: # down + if offset > 0: + B_part = B.crop((0, 0, w, offset)) + if offset < h: + B_part = B_part.resize((w, offset), Image.LANCZOS) + canvas.paste(B_part, (0, h - offset)) + + if offset < h: + A_part = A.crop((0, offset, w, h)) + canvas.paste(A_part, (0, 0)) + + if offset > blend_width and offset < h: + blend_mask = Image.new("L", (w, blend_width)) + blend_draw = ImageDraw.Draw(blend_mask) + for y in range(blend_width): + alpha = int(255 * (blend_width - y) / blend_width) + blend_draw.rectangle([0, y, w, y + 1], fill=alpha) + blend_region = Image.composite( + A.crop((0, offset - blend_width, w, offset)), + B.crop((0, 0, w, blend_width)), + blend_mask + ) + canvas.paste(blend_region, (0, offset - blend_width)) + return canvas + + # ---------------------------- + # 过渡:zoom(优化版 - 使用缓动和更平滑的缩放) + # ---------------------------- + def _zoom(self, A, B, t, mode): + w, h = A.size + # 使用缓动函数让缩放更自然 + t_smooth = self._ease_out_cubic(t) + + if mode == "in": + # B 从 70% 缩放到 100%,使用更平滑的曲线 + scale_start = 0.7 + scale = scale_start + (1.0 - scale_start) * t_smooth + + # 同时进行淡入效果 + blend_t = self._smooth_step(t) + + Bw = max(1, int(w * scale)) + Bh = max(1, int(h * scale)) + B_resized = B.resize((Bw, Bh), Image.LANCZOS) + + # 创建带透明度的混合 + canvas = A.copy() + paste_x = (w - Bw) // 2 + paste_y = (h - Bh) // 2 + + # 使用混合而不是直接粘贴,让过渡更平滑 + B_layer = Image.new("RGB", (w, h)) + B_layer.paste(B_resized, (paste_x, paste_y)) + return Image.blend(A, B_layer, blend_t) + else: + # zoom_out: A 缩小并混合 B,使用更平滑的曲线 + scale_start = 1.0 + scale_end = 0.6 + scale = scale_start - (scale_start - scale_end) * t_smooth + + blend_t = self._smooth_step(t) + + Aw = max(1, int(w * scale)) + Ah = max(1, int(h * scale)) + A_resized = A.resize((Aw, Ah), Image.LANCZOS) + + # 将缩小的A放在B上,然后混合 + A_layer = Image.new("RGB", (w, h)) + paste_x = (w - Aw) // 2 + paste_y = (h - Ah) // 2 + A_layer.paste(A_resized, (paste_x, paste_y)) + + return Image.blend(A_layer, B, blend_t) + + # ---------------------------- + # 过渡:轻微模糊以软化衔接(优化版) + # ---------------------------- + def _soft_blur_blend(self, A, B, t): + # 使用缓动函数 + t_smooth = self._ease_in_out_quad(t) + + # 给 A 应用逐步增强的模糊,然后 blend 到 B,效果柔和 + # 使用更平滑的模糊曲线 + max_blur = 15 + radius = 1 + int(max_blur * t_smooth * t_smooth) # 二次曲线让模糊更自然 + + Ab = A.filter(ImageFilter.GaussianBlur(radius=min(radius, max_blur))) + + # 同时给B也添加轻微模糊,让过渡更平滑 + if t_smooth > 0.3: + B_blur_radius = int(3 * (t_smooth - 0.3) / 0.7) + if B_blur_radius > 0: + B = B.filter(ImageFilter.GaussianBlur(radius=min(B_blur_radius, 5))) + + return Image.blend(Ab, B, t_smooth) + + # ---------------------------- + # 过渡:旋转过渡(新增 - 让过渡更丰富) + # ---------------------------- + def _rotate_transition(self, A, B, t, direction="clockwise"): + w, h = A.size + # 使用缓动函数 + t_smooth = self._ease_in_out_cubic(t) + + # 旋转角度:从0度到45度(更温和的旋转) + max_angle = 45 + if direction == "clockwise": + angle = max_angle * t_smooth + else: # counterclockwise + angle = -max_angle * t_smooth + + # A旋转并淡出 + A_rotated = A.rotate(angle, expand=False, resample=Image.BICUBIC, fillcolor=(0, 0, 0)) + A_alpha = int(255 * (1 - t_smooth)) + + # B旋转并淡入(反向旋转) + B_angle = -angle if direction == "clockwise" else angle + B_rotated = B.rotate(B_angle, expand=False, resample=Image.BICUBIC, fillcolor=(0, 0, 0)) + B_alpha = int(255 * t_smooth) + + # 创建带透明度的图层 + A_rgba = A_rotated.convert("RGBA") + A_alpha_channel = Image.new("L", (w, h), A_alpha) + A_layer = Image.merge("RGBA", (*A_rgba.split()[:3], A_alpha_channel)) + + B_rgba = B_rotated.convert("RGBA") + B_alpha_channel = Image.new("L", (w, h), B_alpha) + B_layer = Image.merge("RGBA", (*B_rgba.split()[:3], B_alpha_channel)) + + # 合成 + canvas = Image.new("RGBA", (w, h), (0, 0, 0, 255)) + canvas = Image.alpha_composite(canvas, A_layer) + canvas = Image.alpha_composite(canvas, B_layer) + + return canvas.convert("RGB") + + # ---------------------------- + # 统一的过渡入口(返回 PIL.Image) + # ---------------------------- + def apply_transition_frame(self, pathA, pathB, t, transition_type): + A, B = self._open_and_match(pathA, pathB) + + # 若为 random,则在具体候选中随机选择 + if transition_type == "random": + candidates = [ + "crossfade", "left_wipe", "right_wipe", #"up_wipe", "down_wipe", + "slide_left", "slide_right", #"slide_up", "slide_down", + "zoom_in", "zoom_out", "soft_blur" + ] + transition_type = random.choice(candidates) + + if transition_type == "crossfade": + return self._crossfade(A, B, t) + if transition_type in ("left_wipe", "right_wipe", "up_wipe", "down_wipe"): + dir_map = { + "left_wipe": "left", "right_wipe": "right", + "up_wipe": "up", "down_wipe": "down" + } + return self._wipe(A, B, t, dir_map[transition_type]) + if transition_type.startswith("slide_"): + dir_map = { + "slide_left": "left", "slide_right": "right", + "slide_up": "up", "slide_down": "down" + } + return self._slide(A, B, t, dir_map[transition_type]) + if transition_type == "zoom_in": + return self._zoom(A, B, t, "in") + if transition_type == "zoom_out": + return self._zoom(A, B, t, "out") + if transition_type == "soft_blur": + return self._soft_blur_blend(A, B, t) + + # fallback + return self._crossfade(A, B, t) + + # ---------------------------- + # 音频相关(与原逻辑一致) + # ---------------------------- def get_audio_duration(self, audio_path): - """获取音频时长(秒)""" + cmd = [ + "ffprobe", "-v", "error", + "-show_entries", "format=duration", + "-of", "default=noprint_wrappers=1:nokey=1", + audio_path + ] + result = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) + out = result.stdout.strip() try: - cmd = [ - "ffprobe", - "-v", "error", - "-show_entries", "format=duration", - "-of", "default=noprint_wrappers=1:nokey=1", - audio_path - ] - result = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) - duration = float(result.stdout.strip()) - return duration - except Exception as e: - print(f"Error getting audio duration: {e}") + return float(out) + except: return 0.0 - + def save_audio_to_temp_file(self, audio_input): - """ - 将 ComfyUI 的 AUDIO 对象 ({"waveform": tensor, "sample_rate": int}) - 保存为临时的 .wav 文件供 ffmpeg 使用 - """ - # 解析 AUDIO 对象结构 - waveform = audio_input['waveform'] # Shape通常是 [batch, channels, time] 或 [channels, time] - sample_rate = audio_input['sample_rate'] + waveform = audio_input["waveform"] + sample_rate = audio_input["sample_rate"] - # 确保 waveform 是 CPU tensor if isinstance(waveform, torch.Tensor): waveform = waveform.cpu() - - # 处理维度问题:torchaudio.save 期望 (channels, time) - # 如果是 3维 (1, C, T),需要 squeeze if waveform.dim() == 3: waveform = waveform.squeeze(0) - - # 创建临时文件 - # delete=False 因为我们需要关闭文件后让 ffmpeg 读取它 - temp_audio = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) - temp_audio_path = temp_audio.name - temp_audio.close() # 关闭句柄,让 torchaudio 写入 - try: - torchaudio.save(temp_audio_path, waveform, sample_rate) - print(f"音频对象已缓存至: {temp_audio_path}") - return temp_audio_path - except Exception as e: - if os.path.exists(temp_audio_path): - os.remove(temp_audio_path) - raise ValueError(f"保存临时音频文件失败: {e}") - - def execute_video_synth(self, image_paths, frame_counts, audio, fps, filename_prefix, format="mp4"): - # 0. 预处理音频:将 AUDIO 对象转为临时文件路径 - # audio 参数现在是来自 LoadAudioUpload 的字典 + tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) + tmp_path = tmp.name + tmp.close() + + torchaudio.save(tmp_path, waveform, sample_rate) + return tmp_path + + # ---------------------------- + # 主逻辑:执行合成 + # ---------------------------- + def execute_video_synth( + self, + image_paths, + frame_counts, + audio, + fps, + filename_prefix, + format="mp4", + enable_transition=True, + transition_seconds=0.5, + transition_type="crossfade", + ): + # 1. 保存音频 audio_temp_path = self.save_audio_to_temp_file(audio) - # 1. 解析输入路径和帧数 - paths_list = [p.strip() for p in image_paths.split('\n') if p.strip()] - counts_list = [int(c.strip()) for c in frame_counts.split('\n') if c.strip()] + # 2. 解析路径和帧数 + paths_list = [p.strip() for p in image_paths.split("\n") if p.strip()] + counts_list = [int(c.strip()) for c in frame_counts.split("\n") if c.strip()] if len(paths_list) != len(counts_list): - raise ValueError(f"错误: 路径数量 ({len(paths_list)}) 与 帧数设定数量 ({len(counts_list)}) 不一致。") + raise ValueError("路径数量与帧数数量不一致") - - # 2. 收集所有图片帧 all_frames = [] - print(f"开始处理 {len(paths_list)} 组输入...") - + folder_start_indices = [] + current_index = 0 + for path, count in zip(paths_list, counts_list): - # 去除可能存在的引号 path = path.strip('"').strip("'") - - if not os.path.exists(path): - print(f"警告: 路径不存在,跳过: {path}") - continue + folder_start_indices.append(current_index) if os.path.isdir(path): - # 读取文件夹内图片 valid_exts = ('*.png', '*.jpg', '*.jpeg', '*.bmp', '*.webp') images = [] for ext in valid_exts: images.extend(glob.glob(os.path.join(path, ext))) - - # 按文件名排序确保顺序正确 images.sort() - + if images: - # 截取前 count 张 selected = images[:count] - # 如果文件夹里的图不够 count 数量,用最后一张图补齐 if len(selected) < count: - diff = count - len(selected) - selected.extend([selected[-1]] * diff) - + selected.extend([selected[-1]] * (count - len(selected))) all_frames.extend(selected) - - elif os.path.isfile(path): - # 单张图片重复 count 次 + else: + # 文件夹为空就跳过(不会增加索引) + continue + else: + # 单文件重复 count 次 all_frames.extend([path] * count) + current_index += count + total_frames = len(all_frames) if total_frames == 0: raise ValueError("错误: 未收集到任何图片帧。") - # 3. 计算同步所需的输入帧率 (Input FPS) + print(f"收集到总帧数: {total_frames}") + + # 3. 过渡处理(替换边界帧,保持总帧数不变) + # 收集所有过渡帧临时文件路径,用于后续清理 + transition_temp_files = [] + + if enable_transition and len(folder_start_indices) > 1: + # 过渡帧数(至少 1) + fade_len = max(1, int(round(fps * transition_seconds))) + print(f"[Transition] 启用,类型: {transition_type}, 秒: {transition_seconds}, 帧数: {fade_len}") + + # 注意:我们将用过渡帧替换边界处的原始帧,而不是插入新帧 + # 这样可以保持总帧数不变 + # 原始文件路径保持不变,过渡帧使用临时文件 + for i in range(len(folder_start_indices) - 1): + A_start = folder_start_indices[i] + B_start = folder_start_indices[i + 1] + A_end = B_start - 1 + + # 计算两侧可用帧数(保证过渡不超边界) + available_A = A_end - A_start + 1 + next_folder_start = folder_start_indices[i + 2] if (i + 2) < len(folder_start_indices) else total_frames + available_B = next_folder_start - B_start + + # 先确定实际可替换的帧数 + replace_from_A = min(fade_len // 2, available_A) + replace_from_B = min(fade_len - fade_len // 2, available_B) + actual_replace = replace_from_A + replace_from_B + + if actual_replace <= 1: + continue + + # 若是 random:为当前段随机选择一个过渡类型(更自然) + chosen_type = transition_type + if transition_type == "random": + opts = [ + "crossfade", "left_wipe", "right_wipe", "up_wipe", "down_wipe", + "slide_left", "slide_right", "slide_up", "slide_down", + "zoom_in", "zoom_out", "soft_blur" + ] + chosen_type = random.choice(opts) + + print(f"[Transition] 段 {i}→{i+1} | 类型: {chosen_type} | 使用帧数: {actual_replace} (A:{replace_from_A}, B:{replace_from_B})") + + # 生成 transition 帧 + # 关键优化:在过渡过程中,B画面应该使用B段中对应位置的帧,而不是只用B的第一帧 + # 这样B画面在过渡过程中会动态变化,而不是静态的 + # 例如在"向左推"效果中,右侧的B画面会随着时间播放,而不是只显示B的第一帧 + transition_paths = [] + for k in range(actual_replace): + # t值从0到1均匀分布,确保整个过渡过程都有变化 + t = k / (actual_replace - 1) if actual_replace > 1 else 0.0 + + # 根据过渡进度,选择A段和B段中对应位置的帧 + # A段:从末尾往前取,让A画面在过渡过程中也动态变化 + # 让A画面从A段的末尾开始,随着过渡进度逐渐使用A段更早的帧 + # 这样在"向左推"等效果中,左侧的A画面会随着时间播放,而不是静态的 + A_frame_idx = A_end - k + # 确保不超出A段范围 + A_frame_idx = max(A_frame_idx, A_start) + + # B段:从开始往后取,让B画面在整个过渡过程中动态变化 + # 让B画面从B的第一帧开始,随着过渡进度逐渐使用B段后续的帧 + # 这样在"向左推"等效果中,右侧的B画面会随着时间播放,而不是静态的 + B_frame_idx = B_start + k + # 确保不超出B段范围 + B_frame_idx = min(B_frame_idx, B_start + available_B - 1) + + A_path = all_frames[A_frame_idx] + B_path = all_frames[B_frame_idx] + + # 读取原始文件生成过渡帧(不修改原始文件) + img = self.apply_transition_frame(A_path, B_path, t, chosen_type) + # 将过渡帧保存到临时文件(不会影响原始文件) + tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False) + tmp_path = tmp.name + tmp.close() + img.save(tmp_path) + transition_paths.append(tmp_path) + # 记录临时文件路径,用于后续清理 + transition_temp_files.append(tmp_path) + + # 替换 A 段最后的帧(使用过渡帧的前半部分,从 A 到中间) + replace_A_start = A_end - replace_from_A + 1 + for idx in range(replace_from_A): + # 使用对应的过渡帧索引,确保 t 值从 0 开始递增 + trans_path = transition_paths[idx] + all_frames[replace_A_start + idx] = trans_path + + # 替换 B 段前面的帧(使用过渡帧的后半部分,从中间到 B) + for idx in range(replace_from_B): + # 使用过渡帧的后半部分,索引从 replace_from_A 开始 + trans_path = transition_paths[replace_from_A + idx] + all_frames[B_start + idx] = trans_path + + # 关键修复:删除B段中已经被过渡帧替换的原始帧,避免重复 + # 由于过渡帧已经替换了B段的前 replace_from_B 帧 + # 需要从 all_frames 中删除B段中从 replace_from_B 开始的原始帧 + # 删除从 B_start + replace_from_B 开始的 replace_from_B 个帧 + del_start = B_start + replace_from_B + del_end = del_start + replace_from_B + if del_end <= len(all_frames): + # 删除B段中已经被替换的原始帧 + del all_frames[del_start:del_end] + # 更新总帧数 + total_frames -= replace_from_B + # 更新后续段的起始索引 + for j in range(i + 2, len(folder_start_indices)): + folder_start_indices[j] -= replace_from_B + + print(f"[Transition] 替换: A段最后{replace_from_A}帧, B段前{replace_from_B}帧, 删除B段重复帧{replace_from_B}帧, 总帧数: {total_frames}") + + else: + print("[Transition] 未启用或无多段,跳过过渡处理") + + # 4. 将帧复制到临时目录并用 ffmpeg 合成(与之前逻辑一致) + temp_dir = os.path.join(folder_paths.get_temp_directory(), "debug_frames_video_synth") + if os.path.exists(temp_dir): + shutil.rmtree(temp_dir) + os.makedirs(temp_dir, exist_ok=True) + + for i, img_path in enumerate(all_frames): + shutil.copy(img_path, os.path.join(temp_dir, f"{i:05d}.png")) + + # 计算 input_fps,使得帧按音频时长播放 audio_duration = self.get_audio_duration(audio_temp_path) if audio_duration <= 0: raise ValueError("错误: 音频时长无效。") - - # 计算逻辑:为了让 all_frames 刚好播完等于 audio_duration,输入流速度必须是 input_fps - input_fps = total_frames / audio_duration - print(f"统计: 总帧数 {total_frames} | 音频时长 {audio_duration}s | 计算输入流速: {input_fps:.4f} fps") - print(f"设置: 目标输出视频帧率: {fps} fps") + input_fps = len(all_frames) / audio_duration + + print(f"最终帧数: {len(all_frames)}, 音频时长: {audio_duration}s, input_fps: {input_fps:.4f}") - # 4. 准备输出文件 output_dir = folder_paths.get_output_directory() output_filename = f"{filename_prefix}.{format}" output_file_path = os.path.join(output_dir, output_filename) - - # 避免文件名冲突 counter = 1 while os.path.exists(output_file_path): output_filename = f"{filename_prefix}_{counter}.{format}" output_file_path = os.path.join(output_dir, output_filename) counter += 1 - # 5. 临时目录处理与合成 - # 1. 不使用临时目录,而是使用 ComfyUI 根目录下的 temp/debug_frames 文件夹 - # 获取 ComfyUI 的 temp 目录 - comfy_temp_dir = folder_paths.get_temp_directory() - # 创建一个固定的子文件夹,每次运行前清空(或者不清空看你需要) - temp_dir = os.path.join(comfy_temp_dir, "debug_frames_video_synth") - - # 如果目录存在,先清空旧文件(防止混淆),如果不存在则创建 - if os.path.exists(temp_dir): - shutil.rmtree(temp_dir) - os.makedirs(temp_dir) + ffmpeg_cmd = [ + "ffmpeg", + "-y", + "-framerate", str(input_fps), + "-i", os.path.join(temp_dir, "%05d.png"), + "-i", audio_temp_path, + "-r", str(fps), + "-c:v", "libx264", + "-pix_fmt", "yuv420p", + "-c:a", "copy", + "-shortest", + output_file_path + ] - print(f"★ 中间帧图片已保存在: {temp_dir}") + subprocess.run(ffmpeg_cmd, check=True) - # 去掉 'with tempfile...' 的缩进,直接执行 + # 清理临时文件 + # 清理临时音频 try: - print("正在准备帧序列...") - for i, img_path in enumerate(all_frames): - temp_img_name = f"{i:05d}.png" - shutil.copy(img_path, os.path.join(temp_dir, temp_img_name)) + if os.path.exists(audio_temp_path): + os.remove(audio_temp_path) + except Exception: + pass + + # 清理过渡帧临时文件(确保不留下垃圾文件,原始文件路径保持不变) + for temp_file in transition_temp_files: + try: + if os.path.exists(temp_file): + os.remove(temp_file) + except Exception: + pass - print("开始执行 FFmpeg 合成...") - - ffmpeg_cmd = [ - "ffmpeg", - "-y", - "-framerate", str(input_fps), - "-i", os.path.join(temp_dir, "%05d.png"), - "-i", audio_temp_path, - "-r", str(fps), - "-c:v", "libx264", - "-pix_fmt", "yuv420p", - "-c:a", "copy", - "-shortest", - output_file_path - ] - - subprocess.run(ffmpeg_cmd, check=True) - - except Exception as e: - raise e - finally: - # 6. 清理临时音频文件 - # 无论成功还是失败,都要把 save_audio_to_temp_file 生成的文件删掉 - if audio_temp_path and os.path.exists(audio_temp_path): - try: - os.remove(audio_temp_path) - print(f"已清理临时音频文件: {audio_temp_path}") - except Exception as e: - print(f"清理临时音频文件失败: {e}") - - print(f"视频生成完毕: {output_file_path}") + print(f"视频生成成功: {output_file_path}") return (output_file_path,) + +class WanVideoImageToVideoInfiniteTalkFunCamera: + @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": 25, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}), + "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', + '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"}), + "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), + "clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}), + "control_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "Control signal for the Fun -model"}), + "fun_or_fl2v_model": ("BOOLEAN", {"default": True, "tooltip": "Enable when using official FLF2V or Fun model"}), + "mode": ([ + "auto", + "multitalk", + "infinitetalk" + ], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"}), + "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, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, control_embeds = None, fun_or_fl2v_model=False, mode="multitalk", output_path=""): + + H = height + W = width + VAE_STRIDE = (4, 8, 8) + + 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) // VAE_STRIDE[0] + 1, + height // VAE_STRIDE[1], + width // VAE_STRIDE[2]) + + 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) + + 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, + "target_h": H, + "target_w": W, + "lat_h": H, + "lat_w": W, + "control_embeds": control_embeds["control_embeds"] if control_embeds is not None else None, + "tiled_vae": tiled_vae, + "force_offload": force_offload, + "fun_or_fl2v_model": fun_or_fl2v_model, + "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": mode, + "output_path": output_path + } + + return (image_embeds, output_path) + + +from comfy import model_management as mm +import os, gc, math +from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device) + +VAE_STRIDE = (4, 8, 8) +PATCH_SIZE = (1, 2, 2) + +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + +class WanVideoImageToVideoEncodeForInfiniteTalkFunCamera: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), + "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), + "noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation, helpful for I2V where some noise can add motion and give sharper results"}), + "start_latent_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional latent multiplier, helpful for I2V where lower values allow for more motion"}), + "end_latent_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional latent multiplier, helpful for I2V where lower values allow for more motion"}), + "force_offload": ("BOOLEAN", {"default": True}), + }, + "optional": { + "vae": ("WANVAE",), + "clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}), + "start_image": ("IMAGE", {"tooltip": "Image to encode"}), + "end_image": ("IMAGE", {"tooltip": "end frame"}), + "control_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "Control signal for the Fun -model"}), + "fun_or_fl2v_model": ("BOOLEAN", {"default": True, "tooltip": "Enable when using official FLF2V or Fun model"}), + "temporal_mask": ("MASK", {"tooltip": "mask"}), + "extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}), + "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), + "add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}), + "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",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, width, height, num_frames, force_offload, noise_aug_strength, + start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, + temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, output_path=""): + + if vae is None: + raise ValueError("VAE is required for image encoding.") + H = height + W = width + + lat_h = H // vae.upsampling_factor + lat_w = W // vae.upsampling_factor + + num_frames = ((num_frames - 1) // 4) * 4 + 1 + two_ref_images = start_image is not None and end_image is not None + + if start_image is None and end_image is not None: + fun_or_fl2v_model = True # end image alone only works with this option + + base_frames = num_frames + (1 if two_ref_images and not fun_or_fl2v_model else 0) + if temporal_mask is None: + mask = torch.zeros(1, base_frames, lat_h, lat_w, device=device, dtype=vae.dtype) + if start_image is not None: + mask[:, 0:start_image.shape[0]] = 1 # First frame + if end_image is not None: + mask[:, -end_image.shape[0]:] = 1 # End frame if exists + else: + mask = common_upscale(temporal_mask.unsqueeze(1).to(device), lat_w, lat_h, "nearest", "disabled").squeeze(1) + if mask.shape[0] > base_frames: + mask = mask[:base_frames] + elif mask.shape[0] < base_frames: + mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)]) + mask = mask.unsqueeze(0).to(device, vae.dtype) + + # Repeat first frame and optionally end frame + start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W + if end_image is not None and not fun_or_fl2v_model: + end_mask_repeated = torch.repeat_interleave(mask[:, -1:], repeats=4, dim=1) # T, C, H, W + mask = torch.cat([start_mask_repeated, mask[:, 1:-1], end_mask_repeated], dim=1) + else: + mask = torch.cat([start_mask_repeated, mask[:, 1:]], dim=1) + + # Reshape mask into groups of 4 frames + mask = mask.view(1, mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W + mask = mask.movedim(1, 2)[0]# C, T, H, W + + # Resize and rearrange the input image dimensions + if start_image is not None: + start_image = start_image[..., :3] + if start_image.shape[1] != H or start_image.shape[2] != W: + resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_start_image = start_image.permute(3, 0, 1, 2) # C, T, H, W + resized_start_image = resized_start_image * 2 - 1 + if noise_aug_strength > 0.0: + resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength) + + if end_image is not None: + end_image = end_image[..., :3] + if end_image.shape[1] != H or end_image.shape[2] != W: + resized_end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_end_image = end_image.permute(3, 0, 1, 2) # C, T, H, W + resized_end_image = resized_end_image * 2 - 1 + if noise_aug_strength > 0.0: + resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength) + + # Concatenate image with zero frames and encode + if temporal_mask is None: + if start_image is not None and end_image is None: + zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1) + del resized_start_image, zero_frames + elif start_image is None and end_image is not None: + zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) + del zero_frames + elif start_image is None and end_image is None: + concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype) + else: + if fun_or_fl2v_model: + zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype) + else: + zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) + del resized_start_image, zero_frames + else: + temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1) + concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype) + del resized_start_image, temporal_mask + + mm.soft_empty_cache() + gc.collect() + + vae.to(device) + y = vae.encode([concatenated], device, end_=(end_image is not None and not fun_or_fl2v_model),tiled=tiled_vae)[0] + del concatenated + + has_ref = False + if extra_latents is not None: + samples = extra_latents["samples"].squeeze(0) + y = torch.cat([samples, y], dim=1) + mask = torch.cat([torch.ones_like(mask[:, 0:samples.shape[1]]), mask], dim=1) + num_frames += samples.shape[1] * 4 + has_ref = True + y[:, :1] *= start_latent_strength + y[:, -1:] *= end_latent_strength + + # Calculate maximum sequence length + patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2]) + frames_per_stride = (num_frames - 1) // 4 + (2 if end_image is not None and not fun_or_fl2v_model else 1) + max_seq_len = frames_per_stride * patches_per_frame + + if add_cond_latents is not None: + add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device) + + if force_offload: + vae.model.to(offload_device) + mm.soft_empty_cache() + gc.collect() + if output_path: + timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + output_path = os.path.join(output_path, f"{timestamp}_infinitetalk_output") + os.makedirs(output_path, exist_ok=True) + image_embeds = { + "image_embeds": y.cpu(), + "clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None, + "negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None, + "max_seq_len": max_seq_len, + "num_frames": num_frames, + "lat_h": lat_h, + "lat_w": lat_w, + "control_embeds": control_embeds["control_embeds"] if control_embeds is not None else None, + "end_image": resized_end_image if end_image is not None else None, + "fun_or_fl2v_model": fun_or_fl2v_model, + "has_ref": has_ref, + "add_cond_latents": add_cond_latents, + "mask": mask.cpu(), + "output_path": output_path + } + + return (image_embeds,) + NODE_CLASS_MAPPINGS = { "InfiniteTalkMultiImage": InfiniteTalkMultiImage, + "WanVideoImageToVideoInfiniteTalkFunCamera": WanVideoImageToVideoInfiniteTalkFunCamera, + "WanVideoImageToVideoEncodeForInfiniteTalkFunCamera": WanVideoImageToVideoEncodeForInfiniteTalkFunCamera, + "MakeBatchFromIntList": MakeBatchFromIntList, "GetIntByIndex": GetIntByIndex, "GetFloatByIndex": GetFloatByIndex, @@ -764,6 +1576,7 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "InfiniteTalkMultiImage": "InfiniteTalkMultiImage", + "WanVideoImageToVideoEncodeForInfiniteTalkFunCamera": "WanVideoImageToVideoEncodeForInfiniteTalkFunCamera", "MakeBatchFromIntList": "MakeBatchFromIntList", "GetIntByIndex": "GetIntByIndex", "GetFloatByIndex": "GetFloatByIndex", diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..5b155c6 --- /dev/null +++ b/utils.py @@ -0,0 +1,611 @@ +import importlib.metadata +import torch +import logging +import math +from tqdm import tqdm +from pathlib import Path +import os +import types, collections +from comfy.utils import ProgressBar, copy_to_param, set_attr_param +from comfy.model_patcher import get_key_weight, string_to_seed +from comfy.lora import calculate_weight +from comfy.model_management import cast_to_device +from comfy.float import stochastic_rounding +import folder_paths +logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') +log = logging.getLogger(__name__) + +def check_device_same(first_device, second_device): + if first_device.type != second_device.type: + return False + + if first_device.type == "cuda" and first_device.index is None: + first_device = torch.device("cuda", index=0) + + if second_device.type == "cuda" and second_device.index is None: + second_device = torch.device("cuda", index=0) + + return first_device == second_device + +# simplified version of the accelerate function https://github.com/huggingface/accelerate/blob/main/src/accelerate/utils/modeling.py +def set_module_tensor_to_device(module, tensor_name, device, value=None, dtype=None): + """ + A helper function to set a given tensor (parameter of buffer) of a module on a specific device (note that doing + `param.to(device)` creates a new tensor not linked to the parameter, which is why we need this function). + + Args: + module (`torch.nn.Module`): + The module in which the tensor we want to move lives. + tensor_name (`str`): + The full name of the parameter/buffer. + device (`int`, `str` or `torch.device`): + The device on which to set the tensor. + value (`torch.Tensor`, *optional*): + The value of the tensor (useful when going from the meta device to any other device). + dtype (`torch.dtype`, *optional*): + If passed along the value of the parameter will be cast to this `dtype`. Otherwise, `value` will be cast to + the dtype of the existing parameter in the model. + """ + # Recurse if needed + if "." in tensor_name: + splits = tensor_name.split(".") + for split in splits[:-1]: + new_module = getattr(module, split) + if new_module is None: + raise ValueError(f"{module} has no attribute {split}.") + module = new_module + tensor_name = splits[-1] + + if tensor_name not in module._parameters and tensor_name not in module._buffers: + raise ValueError(f"{module} does not have a parameter or a buffer named {tensor_name}.") + is_buffer = tensor_name in module._buffers + old_value = getattr(module, tensor_name) + + if old_value.device == torch.device("meta") and device not in ["meta", torch.device("meta")] and value is None: + raise ValueError(f"{tensor_name} is on the meta device, we need a `value` to put in on {device}.") + + param = module._parameters[tensor_name] if tensor_name in module._parameters else None + param_cls = type(param) + + if value is not None: + if dtype is None: + value = value.to(old_value.dtype) + elif not str(value.dtype).startswith(("torch.uint", "torch.int", "torch.bool")): + value = value.to(dtype) + + device_quantization = None + with torch.no_grad(): + if value is None: + new_value = old_value.to(device) + if dtype is not None and device in ["meta", torch.device("meta")]: + if not str(old_value.dtype).startswith(("torch.uint", "torch.int", "torch.bool")): + new_value = new_value.to(dtype) + + if not is_buffer: + module._parameters[tensor_name] = param_cls(new_value, requires_grad=old_value.requires_grad) + elif isinstance(value, torch.Tensor): + new_value = value.to(device) + else: + new_value = torch.tensor(value, device=device) + if device_quantization is not None: + device = device_quantization + if is_buffer: + module._buffers[tensor_name] = new_value + 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).to(device) + module._parameters[tensor_name] = new_value + + #if device != "cpu": + # mm.soft_empty_cache() + +def check_diffusers_version(): + try: + version = importlib.metadata.version('diffusers') + required_version = '0.31.0' + if version < required_version: + raise AssertionError(f"diffusers version {version} is installed, but version {required_version} or higher is required.") + except importlib.metadata.PackageNotFoundError: + raise AssertionError("diffusers is not installed.") + +def print_memory(device): + memory = torch.cuda.memory_allocated(device) / 1024**3 + max_memory = torch.cuda.max_memory_allocated(device) / 1024**3 + max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3 + log.info(f"Allocated memory: {memory=:.3f} GB") + log.info(f"Max allocated memory: {max_memory=:.3f} GB") + log.info(f"Max reserved memory: {max_reserved=:.3f} GB") + #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) + #log.info(f"Memory Summary:\n{memory_summary}") + +def get_module_memory_mb(module): + memory = 0 + for param in module.parameters(): + if param.data is not None: + memory += param.nelement() * param.element_size() + return memory / (1024 * 1024) # Convert to MB + +def get_tensor_memory(tensor): + memory_bytes = tensor.element_size() * tensor.nelement() + return f"{memory_bytes / (1024 * 1024):.2f} MB" + +def patch_weight_to_device(self, key, device_to=None, inplace_update=False, backup_keys=False, scale_weight=None): + if key not in self.patches: + return + + weight, set_func, convert_func = get_key_weight(self.model, key) + inplace_update = self.weight_inplace_update or inplace_update + + if backup_keys and key not in self.backup: + self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update) + + if device_to is not None: + temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True) + else: + temp_weight = weight.to(torch.float32, copy=True) + if convert_func is not None: + temp_weight = convert_func(temp_weight, inplace=True) + + if scale_weight is not None: + temp_weight = temp_weight * scale_weight.to(temp_weight.device, temp_weight.dtype) + + out_weight = calculate_weight(self.patches[key], temp_weight, key) + + if set_func is None: + out_weight = stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key)) + if inplace_update: + copy_to_param(self.model, key, out_weight) + else: + set_attr_param(self.model, key, out_weight) + else: + set_func(out_weight, inplace_update=inplace_update, seed=string_to_seed(key)) + +def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None, + base_dtype=None, state_dict=None, low_mem_load=False, control_lora=False, scale_weights={}): + model.patch_weight_to_device = types.MethodType(patch_weight_to_device, model) + to_load = [] + for n, m in model.model.named_modules(): + params = [] + skip = False + for name, param in m.named_parameters(recurse=False): + params.append(name) + for name, param in m.named_parameters(recurse=True): + if name not in params: + skip = True # skip random weights in non leaf modules + break + if not skip and (hasattr(m, "comfy_cast_weights") or len(params) > 0): + to_load.append((n, m, params)) + + to_load.sort(reverse=True) + cnt = 0 + pbar = ProgressBar(len(to_load)) + for x in tqdm(to_load, desc="Loading model and applying LoRA weights:", leave=True): + name = x[0] + m = x[1] + params = x[2] + if hasattr(m, "comfy_patched_weights"): + if m.comfy_patched_weights == True: + continue + for param in params: + name = name.replace("._orig_mod.", ".") # torch compiled modules have this prefix + if low_mem_load: + dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype + if "patch_embedding" in name: + dtype_to_use = torch.float32 + key = f"{name.replace('diffusion_model.', '')}.{param}" + try: + set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key]) + except: + continue + key = f"{name}.{param}" + if scale_weights is not None: + scale_key = key.replace("weight", "scale_weight").replace("diffusion_model.", "") if "weight" in key else None + if low_mem_load: + model.patch_weight_to_device(f"{name}.{param}", device_to=device_to, inplace_update=True, backup_keys=control_lora, scale_weight=scale_weights.get(scale_key, None)) + else: + model.patch_weight_to_device(f"{name}.{param}", device_to=device_to, backup_keys=control_lora, scale_weight=scale_weights.get(scale_key, None)) + if device_to != transformer_load_device: + set_module_tensor_to_device(m, param, device=transformer_load_device) + if low_mem_load: + try: + set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key]) + except: + continue + m.comfy_patched_weights = True + cnt += 1 + if cnt % 100 == 0: + pbar.update(100) + + + # After LoRA patching, scale weights that have scale_weight but are NOT LoRA patched + if len(scale_weights) > 0 and not getattr(model, "scale_weights_applied", False): + for name, param in model.model.diffusion_model.named_parameters(): + scale_key = name.replace("weight", "scale_weight").replace("diffusion_model.", "") if "weight" in name else None + full_param_name = f"diffusion_model.{name}" + if scale_key and scale_key in scale_weights and full_param_name not in model.patches: + scale = scale_weights[scale_key] + param_fp32 = param.to(torch.float32) + param_fp32.mul_(scale.to(param.device, torch.float32)) + param.copy_(param_fp32.to(param.dtype)) + model.scale_weights_applied = True + + model.current_weight_patches_uuid = model.patches_uuid + if low_mem_load: + for name, param in model.model.diffusion_model.named_parameters(): + if param.device != transformer_load_device: + dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype + if "patch_embedding" in name: + dtype_to_use = torch.float32 + try: + set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name]) + except: + continue + return model + + +# from https://github.com/cubiq/ComfyUI_IPAdapter_plus/blob/9d076a3df0d2763cef5510ec5ab807f6632c39f5/utils.py#L181 +def split_tiles(embeds, num_split): + _, H, W, _ = embeds.shape + out = [] + for x in embeds: + x = x.unsqueeze(0) + h, w = H // num_split, W // num_split + x_split = torch.cat([x[:, i*h:(i+1)*h, j*w:(j+1)*w, :] for i in range(num_split) for j in range(num_split)], dim=0) + out.append(x_split) + + x_split = torch.stack(out, dim=0) + + return x_split + +def merge_hiddenstates(x, tiles): + chunk_size = tiles*tiles + x = x.split(chunk_size) + + out = [] + for embeds in x: + num_tiles = embeds.shape[0] + tile_size = int((embeds.shape[1]-1) ** 0.5) + grid_size = int(num_tiles ** 0.5) + + # Extract class tokens + class_tokens = embeds[:, 0, :] # Save class tokens: [num_tiles, embeds[-1]] + avg_class_token = class_tokens.mean(dim=0, keepdim=True).unsqueeze(0) # Average token, shape: [1, 1, embeds[-1]] + + patch_embeds = embeds[:, 1:, :] # Shape: [num_tiles, tile_size^2, embeds[-1]] + reshaped = patch_embeds.reshape(grid_size, grid_size, tile_size, tile_size, embeds.shape[-1]) + + merged = torch.cat([torch.cat([reshaped[i, j] for j in range(grid_size)], dim=1) + for i in range(grid_size)], dim=0) + + merged = merged.unsqueeze(0) # Shape: [1, grid_size*tile_size, grid_size*tile_size, embeds[-1]] + + # Pool to original size + pooled = torch.nn.functional.adaptive_avg_pool2d(merged.permute(0, 3, 1, 2), (tile_size, tile_size)).permute(0, 2, 3, 1) + flattened = pooled.reshape(1, tile_size*tile_size, embeds.shape[-1]) + + # Add back the class token + with_class = torch.cat([avg_class_token, flattened], dim=1) # Shape: original shape + out.append(with_class) + + out = torch.cat(out, dim=0) + + return out + +from comfy.clip_vision import clip_preprocess, ClipVisionModel + +def clip_encode_image_tiled(clip_vision, image, tiles=1, ratio=1.0): + embeds = encode_image_(clip_vision, image) + tiles = min(tiles, 16) + + if tiles > 1: + # split in tiles + image_split = split_tiles(image, tiles) + + # get the embeds for each tile + embeds_split = {} + for i in image_split: + encoded = encode_image_(clip_vision, i) + if not hasattr(embeds_split, "last_hidden_state"): + embeds_split["last_hidden_state"] = encoded + else: + embeds_split["last_hidden_state"] = torch.cat(embeds_split["last_hidden_state"], encoded, dim=0) + + embeds_split['last_hidden_state'] = merge_hiddenstates(embeds_split['last_hidden_state'], tiles) + + if embeds.shape[0] > 1: # if we have more than one image we need to average the embeddings for consistency + embeds = embeds * ratio + embeds_split['last_hidden_state']*(1-ratio) + else: # otherwise we can concatenate them, they can be averaged later + embeds = torch.cat([embeds * ratio, embeds_split['last_hidden_state']]) + + return embeds + +def encode_image_(clip_vision, image): + if isinstance(clip_vision, ClipVisionModel): + out = clip_vision.encode_image(image).last_hidden_state + else: + pixel_values = clip_preprocess(image, size=224, crop=True).float() + out = clip_vision.visual(pixel_values) + + return out + +# Code based on https://github.com/WikiChao/FreSca (MIT License) +import torch +import torch.fft as fft + +def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20): + """ + Apply frequency-dependent scaling to an image tensor using Fourier transforms. + + Parameters: + x: Input tensor of shape (B, C, H, W) + scale_low: Scaling factor for low-frequency components (default: 1.0) + scale_high: Scaling factor for high-frequency components (default: 1.5) + freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20) + + Returns: + x_filtered: Filtered version of x in spatial domain with frequency-specific scaling applied. + """ + # Preserve input dtype and device + dtype, device = x.dtype, x.device + + # Convert to float32 for FFT computations + x = x.to(torch.float32) + + # 1) Apply FFT and shift low frequencies to center + x_freq = fft.fftn(x, dim=(-2, -1)) + x_freq = fft.fftshift(x_freq, dim=(-2, -1)) + + # 2) Create a mask to scale frequencies differently + C, B, H, W = x_freq.shape + crow, ccol = H // 2, W // 2 + + # Initialize mask with high-frequency scaling factor + mask = torch.ones((C, B, H, W), device=device) * scale_high + + # Apply low-frequency scaling factor to center region + mask[ + ..., + crow - freq_cutoff : crow + freq_cutoff, + ccol - freq_cutoff : ccol + freq_cutoff, + ] = scale_low + + # 3) Apply frequency-specific scaling + x_freq = x_freq * mask + + # 4) Convert back to spatial domain + x_freq = fft.ifftshift(x_freq, dim=(-2, -1)) + x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real + + # 5) Restore original dtype + x_filtered = x_filtered.to(dtype) + + return x_filtered + +def is_image_black(image, threshold=1e-3): + if image.min() < 0: + image = (image + 1) / 2 + return torch.all(image < threshold).item() + +def add_noise_to_reference_video(image, ratio=None): + sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio + image_noise = torch.randn_like(image) * sigma[:, None, None, None] + image_noise = torch.where(image==-1, torch.zeros_like(image), image_noise) + image = image + image_noise + return image + +def optimized_scale(positive_flat, negative_flat): + + # Calculate dot production + dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True) + + # Squared norm of uncondition + squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8 + + # st_star = v_cond^T * v_uncond / ||v_uncond||^2 + st_star = dot_product / squared_norm + + return st_star + +def find_closest_valid_dim(fixed_dim, var_dim, block_size): + for delta in range(1, 17): + for sign in [-1, 1]: + candidate = var_dim + sign * delta + if candidate > 0 and ((fixed_dim * candidate) // 4) % block_size == 0: + return candidate + return var_dim + + # Radial attention setup +def setup_radial_attention(transformer, transformer_options, latent, seq_len, latent_video_length, context_options=None): + if context_options is not None: + context_frames = (context_options["context_frames"] - 1) // 4 + 1 + + dense_timesteps = transformer_options.get("dense_timesteps", 1) + dense_blocks = transformer_options.get("dense_blocks", 1) + dense_vace_blocks = transformer_options.get("dense_vace_blocks", 1) + decay_factor = transformer_options.get("decay_factor", 0.2) + dense_attention_mode = transformer_options.get("dense_attention_mode", "sageattn") + block_size = transformer_options.get("block_size", 128) + + # Calculate closest valid latent sizes + if latent.shape[2] % (block_size/8) != 0 or latent.shape[3] % (block_size/8) != 0: + block_div = int(block_size // 8) + closest_h = round(latent.shape[2] / block_div) * block_div + closest_w = round(latent.shape[3] / block_div) * block_div + raise Exception( + f"Radial attention mode only supports image size divisible by block size. " + f"Got {latent.shape[3] * 8}x{latent.shape[2] * 8} with block size {block_size}.\n" + f"Closest valid sizes: {closest_w * 8}x{closest_h * 8} (width x height in pixels)." + ) + tokens_per_frame = (latent.shape[2] * latent.shape[3]) // 4 + if tokens_per_frame % block_size != 0: + closest_latent_h = find_closest_valid_dim(latent.shape[3], latent.shape[2], block_size) + closest_latent_w = find_closest_valid_dim(latent.shape[2], latent.shape[3], block_size) + raise Exception( + f"Radial attention mode requires tokens per frame ((latent_h * latent_w) // 4) to be divisible by block size ({block_size}).\n" + f"Current size in latent space:{latent.shape[3]}x{latent.shape[2]}, pixel space: {latent.shape[3]*8}x{latent.shape[2]*8} tokens_per_frame={tokens_per_frame}.\n" + f"Try adjusting to one of these latent sizes (in pixels):\n" + f" Height: {latent.shape[2]*8} -> {closest_latent_h * 8}\n" + f" Width: {latent.shape[3]*8} -> {closest_latent_w * 8}\n" + f"Or choose another resolution so that (latent_h * latent_w) // 4 is divisible by {block_size}." + ) + + from .wanvideo.radial_attention.attn_mask import MaskMap + for i, block in enumerate(transformer.blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None + if isinstance(dense_blocks, list): + block.dense_block = i in dense_blocks + else: + block.dense_block = i < dense_blocks + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) + block.dense_attention_mode = dense_attention_mode + block.dense_timesteps = dense_timesteps + block.self_attn.decay_factor = decay_factor + if transformer.vace_layers is not None: + for i, block in enumerate(transformer.vace_blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None + if isinstance(dense_vace_blocks, list): + block.dense_block = i in dense_vace_blocks + else: + block.dense_block = i < dense_vace_blocks + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) + block.dense_attention_mode = dense_attention_mode + block.dense_timesteps = dense_timesteps + block.self_attn.decay_factor = decay_factor + + log.info(f"Radial attention mode enabled.") + log.info(f"dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, decay_factor: {decay_factor}") + log.info(f"dense_blocks: {[i for i, block in enumerate(transformer.blocks) if getattr(block, 'dense_block', False)]})") + + + +def list_to_device(tensor_list, device, dtype=None): + """ + Move all tensors in a list to the specified device and optionally cast to dtype. + """ + return [t.to(device, dtype=dtype) if dtype is not None else t.to(device) for t in tensor_list] + +def dict_to_device(tensor_dict, device, dtype=None): + """ + Move all tensors (and tensor lists) in a dict to the specified device and optionally cast to dtype. + Supports values that are tensors or lists of tensors. + """ + result = {} + for k, v in tensor_dict.items(): + if isinstance(v, torch.Tensor): + result[k] = v.to(device, dtype=dtype) if dtype is not None else v.to(device) + elif isinstance(v, list) and all(isinstance(t, torch.Tensor) for t in v): + result[k] = list_to_device(v, device, dtype) + else: + result[k] = v + return result + +def compile_model(transformer, compile_args=None): + if compile_args is None: + return transformer + if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'): + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] + torch._dynamo.config.force_parameter_static_shapes = compile_args["force_parameter_static_shapes"] + try: + if hasattr(torch._dynamo.config, 'allow_unspec_int_on_nn_module'): + torch._dynamo.config.allow_unspec_int_on_nn_module = True + except Exception as e: + log.warning(f"Could not set allow_unspec_int_on_nn_module: {e}") + try: + torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"] + except Exception as e: + log.warning(f"Could not set recompile_limit: {e}") + + if compile_args["compile_transformer_blocks_only"]: + for i, block in enumerate(transformer.blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + transformer.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if transformer.vace_layers is not None: + for i, block in enumerate(transformer.vace_blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + transformer.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + else: + transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + return transformer + +#https://5410tiffany.github.io/tcfg.github.io/ +def tangential_projection(pred_cond: torch.Tensor, pred_uncond: torch.Tensor) -> torch.Tensor: + cond_dtype = pred_cond.dtype + preds = torch.stack([pred_cond, pred_uncond], dim=1).float() + orig_shape = preds.shape[2:] + preds_flat = preds.flatten(2) + U, S, Vh = torch.linalg.svd(preds_flat, full_matrices=False) + Vh_modified = Vh.clone() + Vh_modified[:, 1] = 0 + recon = U @ torch.diag_embed(S) @ Vh_modified + return recon[:, 1].view(pred_uncond.shape).to(cond_dtype) + +#https://arxiv.org/abs/2508.03442 +def get_raag_guidance(noise_pred_cond, noise_pred_uncond, w_max, alpha=1.0, eps=1e-8): + delta = noise_pred_cond - noise_pred_uncond + norm_delta = torch.norm(delta.flatten(1), dim=1, keepdim=True) + norm_uncond = torch.norm(noise_pred_uncond.flatten(1), dim=1, keepdim=True) + ratio = norm_delta / (norm_uncond + eps) + ratio_mean = ratio.mean().item() + adaptive_w = 1.0 + (w_max - 1.0) * math.exp(-alpha * ratio_mean) + return adaptive_w + +def tensor_pingpong_pad(video, target_len): + """ + Pads a video tensor along the frame dimension (dim=2) in a ping-pong fashion. + video: torch.Tensor of shape [B, C, F, H, W] + target_len: desired number of frames + Returns: padded tensor of shape [B, C, target_len, H, W] + """ + in_dims = len(video.shape) + if in_dims == 4: + video = video.unsqueeze(0) + B, C, F, H, W = video.shape + idx = 0 + flip = False + indices = [] + while len(indices) < target_len: + indices.append(idx) + if flip: + idx -= 1 + else: + idx += 1 + if idx == 0 or idx == F - 1: + flip = not flip + indices = indices[:target_len] + padded_video = video[:, :, indices, :, :] + if in_dims == 4: + padded_video = padded_video.squeeze(0) + return padded_video + + +def check_duplicate_nodes(): + """Check ComfyUI custom_nodes directory for duplicate installations""" + custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0]) + current_path = Path(__file__).parent + + wanvideo_dirs = [] + + # Check all directories in custom_nodes + for path in custom_nodes_dir.iterdir(): + if (path.is_dir() and + path != current_path and + 'wanvideo' in path.name.lower() and + 'wrapper' in path.name.lower()): + wanvideo_dirs.append(str(path)) + + return wanvideo_dirs + +#https://github.com/temporalscorerescaling/TSR/ +def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.1): + t = (timestep / 1000) + if t == 0.0: + ratio = k + else: + snr_t = (1 - t)**2 / t**2 + ratio = (snr_t * tsr_sigma**2 + 1) / (snr_t * tsr_sigma**2 / k + 1) + + if not t == 1.0: + model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t) + return model_output