50 Commits
Author SHA1 Message Date
bbaudio-2025 93298fd928 Bump Version 2025-10-31 22:19:34 +08:00
bbaudio-2025 74ed730b91 Add files via upload 2025-10-31 22:18:31 +08:00
bbaudio-2025 ca38b2868f LongVideo add support for wan2.2 vace fun
also add vace strength and seed override for LongVideo
2025-10-31 22:08:51 +08:00
bbaudio-2025 054ab62e9f fix for 2025-10-31 21:28:41 +08:00
bbaudio-2025 3a9686a010 fix for comfyui update 2025-10-31 21:28:03 +08:00
bbaudio-2025 52391ee62a add long_video.json 2025-08-27 13:05:55 +08:00
bbaudio-2025 58700c8099 Create video_upscale.json 2025-08-27 13:05:02 +08:00
bbaudio-2025 07ed89b3cb Delete workflows/tile_control_and_regional_prompt.json 2025-08-27 13:03:00 +08:00
bbaudio-2025 b8b70c2b5c Delete workflows/simple.json 2025-08-27 13:02:48 +08:00
bbaudio-2025 adb48b4e9f Delete workflows/advanced.json 2025-08-27 13:02:35 +08:00
bbaudio-2025 5679597026 Delete workflows/LongVideo.json 2025-08-27 13:02:23 +08:00
bbaudio-2025 3ca3c6d2c0 Bump Version 2025-08-10 22:10:02 +08:00
bbaudio-2025 3bc9abe234 Add files via upload 2025-08-10 22:09:03 +08:00
bbaudio-2025 0f00c49f72 Delete workflows/LongVideoWithRefineInit.json 2025-08-10 22:08:45 +08:00
bbaudio-2025 8801c9b4ea Delete workflows/LongVideoGen.json 2025-08-10 22:08:32 +08:00
bbaudio-2025 d76cd98c98 Update nodes.py
Add NAG options
Add more custom refine options, refact refine code
Move CLIP to main node
2025-08-10 22:07:56 +08:00
bbaudio-2025 99820833b5 Update model.py 2025-08-10 22:03:04 +08:00
bbaudio-2025 0fc8c5fe14 Create model.py 2025-08-10 22:02:38 +08:00
bbaudio-2025 245a5920da Delete nag/wan 2025-08-10 22:02:12 +08:00
bbaudio-2025 d6e7c88881 Add files via upload 2025-08-10 22:01:11 +08:00
bbaudio-2025 3a5678505b Create subfolders 2025-08-10 22:00:31 +08:00
bbaudio-2025 8307159728 Bump Version 2025-08-01 21:54:33 +08:00
bbaudio-2025 172e22d00a add colormatch to vaceupscaler 2025-08-01 21:53:40 +08:00
bbaudio-2025 670875de24 Add model_override and latent_strength_list 2025-07-25 00:00:51 +08:00
bbaudio-2025 fdd2e72d08 fix bug 2025-07-23 13:33:11 +08:00
bbaudio-2025 762bcb2013 fix bug for custom mask 2025-07-22 09:54:15 +08:00
bbaudio-2025 ceb84ee3b2 Update LongVideoWithRefineInit.json 2025-07-22 01:03:43 +08:00
bbaudio-2025 8067f227c4 Bump Version 2025-07-22 00:17:00 +08:00
bbaudio-2025 935ed652a0 Create requirements.txt 2025-07-21 23:49:15 +08:00
bbaudio-2025 38191b6f5e forgot to add colormatch, now is fixed 2025-07-21 23:35:44 +08:00
bbaudio-2025 0f60de4d4b forgot to delete, now it's removed 2025-07-21 21:59:41 +08:00
bbaudio-2025 4991809da2 Update README.md 2025-07-21 21:20:54 +08:00
bbaudio-2025 8ba542d03b Update README.md 2025-07-21 21:20:15 +08:00
bbaudio-2025 22bea2b5f8 Bump Version 2025-07-21 17:51:54 +08:00
bbaudio-2025 93900e1cf6 Add files via upload 2025-07-21 17:51:17 +08:00
bbaudio-2025 f999edee89 Add custom refine init 2025-07-21 17:28:55 +08:00
bbaudio-2025 503e580a85 Add refine_init 2025-07-20 15:15:26 +08:00
bbaudio-2025 60facc895a Add files via upload 2025-07-19 22:38:08 +08:00
bbaudio-2025 923fa30e30 Update nodes.py 2025-07-19 22:36:53 +08:00
bbaudio-2025 798ea62e44 Update pyproject.toml 2025-07-17 20:35:43 +08:00
bbaudio-2025 27e1fb25f3 bug fix 2025-07-17 20:35:18 +08:00
bbaudio-2025 48faf262ec Update README.md 2025-07-16 20:44:44 +08:00
bbaudio-2025 dd649aa05c Version Update 2025-07-16 20:43:31 +08:00
bbaudio-2025 a07081dccc Add "SuperUltimate VACE Long Video" 2025-07-16 20:41:27 +08:00
bbaudio-2025 719bdedfd6 typo 2025-07-15 21:01:20 +08:00
bbaudio-2025 71383e398c add raise error in temporalistgen 2025-07-15 20:59:44 +08:00
bbaudio-2025 2ede0e1f3c Bump Version 2025-07-15 20:30:43 +08:00
bbaudio-2025 5a875855a6 function temporalistgen refactor, bug fixed 2025-07-15 20:29:08 +08:00
bbaudio-2025 f5e7cae1ba fix bug when no control video 2025-07-15 14:33:20 +08:00
bbaudio-2025 9a033d9549 Add files via upload 2025-07-11 23:07:10 +08:00
12 changed files with 7869 additions and 2568 deletions
+25 -2
View File
@@ -4,10 +4,11 @@ SuperUltimateVaceTools, some Comfyui custom nodes for wan2.1 VACE, attempt to im
包含以下插件:
Including following nodes:
- 超究视频放大 | SuperUltimateVaceUpscale
- 超究视频放大 | Super Ultimate Vace Upscale
- 超究长视频 | Super Ultimate VACE Long Video
- 更多功能待添加 | pending to add more...
## SuperUltimateVaceUpscale
## 1. SuperUltimateVaceUpscale
对视频进行分割放大,支持空间分割以及时间分割。
Upscale video by splitting it into tiled areas, supports spatial tiling and temporal tiling.
@@ -33,6 +34,28 @@ To get more controlable result, it is recommended to use reference image and con
After using the reference image, enabling 'crop-ref' will divide the reference image according to the video segmentation plan, and use the segmented reference image for the region reference when denoising each segmented part. If enable 'ref_as_init_frame', the first frame of the video will be directly replaced with the reference image, and this will be used as a reference to guide the following frames denoising of each region.
For control video, VACE supports many control methods, and you can freely try the differences brought by different control methods.
## 2. SuperUltimateVaceLongVideo
利用VACE拼接功能生成长视频,支持多种控制手段,自动修复过渡帧,缓解多轮接续生成带来的视频质量劣化
Generate long length videos with the feature 'temporal extension' of VACE. Support many control methods, automatically refine the crossfade frames, mitigating the quality downgrade from multiple extensions.
### 多轮生成 | Multi-Round generation
你可以在多个`VACE Prompt Combine`节点内编写不同的提示词,指定不同的参考图片,只需要将它们连在一起在连接到`SuperUltimate VACE Long Video`节点。轮次数目没有上限,但由于多轮接续生成花费的时间很长,并且存在不可控的随机性,不建议过长的视频生成。
You can write different prompts within multiple `VACE Prompt Combine` nodes, specify different reference images, and cascade them together and connect to the `SuperUltimate VACE Long Video` node. There is no upper limit to the number of rounds, but due to the long time it takes to generate multiple extensions and the uncontrollable randomness, it is not recommended to generate too long videos.
### 多种控制 ! Multi-Methods control
VACE模型支持许多种控制,包括关键帧、骨骼姿势、深度图、线稿、轨迹动画等等。你可以在一个视频中使用多种不同类型的控制,只需要使用多个`VACE Control Image Combine`节点。但需要注意控制图的帧位不能重复。
如果需要无缝首尾循环视频,只需要为`loopback_crossfade`设置一个大于0的适当数字即可,不需要进行额外的图片控制。
VACE models support many kinds of controls, including keyframes, pose, depth, lineart, trajectory animation, and more. You can use multiple different types of controls in a single video, simply by using multiple `VACE Control Image Combine` nodes. However, you need to be careful that the frame positions of the control image are not duplicated.
If you need a seamless first and last loopback video, just set an appropriate number greater than 0 for `loopback_crossfade`, no additional image controls are needed.
### 缓解多轮接续造成的质量下降 | Mitigating quality degradation due to multiple rounds of generation
VACE可以使用上一轮视频最后几帧作为新一轮生成视频的前几帧,以此实现视频的接续。但一直以来有个问题制约它的应用,那就是随着接续轮次数目增多,生成视频的质量不断下降,出现过饱和、色差、模糊等等问题。
本节点通过为新一轮视频的前几个参考帧进行“修复”的途径,有效缓解了多轮次生成造成的质量下降。但是这么做也会造成一个后果,接续过渡部分的帧的颜色或者亮度会出现不自然变化。
好在副作用并不明显,你也可以自行修改`Custom Refine Option`节点的参数尝试获得更自然的过渡。
VACE can use the last few frames of the previous round of video as the initial few frames of the new round of generation, thus realizing the extension of video. However, there has been a problem that constrains its application, that is, as the number of extent rounds increases, the quality of the generated video decreases, and problems such as oversaturation, chromatic aberration, blurring, etc. occur.
This node effectively mitigates the quality degradation caused by multiple rounds of generation by “refine” the inital few reference frames of a new round of video. However, this also has the consequence that the color or brightness of these frames may change unnaturally.
The side effect is not obvious, but you can also modify the parameters of the `Custom Refine Option` node to try to get a more natural transition.
## 安装 | Install
方法1:在`ComfyUI\custom_nodes`路径下执行命令
Way 1: run following cmd command at the path `ComfyUI\custom_nodes`
+33
View File
@@ -0,0 +1,33 @@
import comfy
from .samplers import KSamplerWithNAG
from .samplers import sample_with_nag as samplers_sample_with_nag
def sample_with_nag(
model, noise, steps, cfg, nag_scale, nag_tau, nag_alpha, nag_sigma_end, sampler_name, scheduler, positive, negative, nag_negative, latent_image, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False, noise_mask=None, sigmas=None, callback=None, disable_pbar=False, seed=None, latent_shapes=None, **kwargs):
sampler = KSamplerWithNAG(model, steps=steps, device=model.load_device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options)
samples = sampler.sample(
noise, positive, negative, nag_negative,
cfg=cfg, nag_scale=nag_scale, nag_tau=nag_tau, nag_alpha=nag_alpha, nag_sigma_end=nag_sigma_end,
latent_image=latent_image,
start_step=start_step, last_step=last_step, force_full_denoise=force_full_denoise,
denoise_mask=noise_mask, sigmas=sigmas, callback=callback, disable_pbar=disable_pbar, seed=seed,
latent_shapes=latent_shapes, **kwargs,
)
samples = samples.to(comfy.model_management.intermediate_device())
return samples
def sample_custom_with_nag(
model, noise, cfg, nag_scale, nag_tau, nag_alpha, nag_sigma_end, sampler, sigmas, positive, negative, nag_negative, latent_image, noise_mask=None, callback=None, disable_pbar=False, seed=None, latent_shapes=None, **kwargs):
samples = samplers_sample_with_nag(
model, noise, positive, negative, nag_negative,
cfg, nag_scale, nag_tau, nag_alpha, nag_sigma_end,
model.load_device, sampler, sigmas,
model_options=model.model_options, latent_image=latent_image, denoise_mask=noise_mask,
callback=callback, disable_pbar=disable_pbar, seed=seed,
latent_shapes=latent_shapes, **kwargs,
)
samples = samples.to(comfy.model_management.intermediate_device())
return samples
+234
View File
@@ -0,0 +1,234 @@
from __future__ import annotations
import copy
from typing import TYPE_CHECKING
import math
if TYPE_CHECKING:
from comfy.model_patcher import ModelPatcher
import torch
from torch._dynamo.eval_frame import OptimizedModule
import torch._dynamo
torch._dynamo.config.suppress_errors = True
from comfy.samplers import (
process_conds,
preprocess_conds_hooks,
cast_to_load_options,
filter_registered_hooks_on_conds,
get_total_hook_groups_in_conds,
CFGGuider,
sampler_object,
KSampler,
)
import comfy.sampler_helpers
import comfy.model_patcher
import comfy.patcher_extension
import comfy.hooks
# from comfy.ldm.flux.model import Flux
# from comfy.ldm.chroma.model import Chroma
# from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel
# from comfy.ldm.modules.diffusionmodules.mmdit import OpenAISignatureMMDITWrapper
from comfy.ldm.wan.model import WanModel, VaceWanModel
# from comfy.ldm.hunyuan_video.model import HunyuanVideo
# from comfy.ldm.hidream.model import HiDreamImageTransformer2DModel
# from .flux.model import NAGFluxSwitch
# from .chroma.model import NAGChromaSwitch
# from .sd.openaimodel import NAGUNetModelSwitch
# from .sd3.mmdit import NAGOpenAISignatureMMDITWrapperSwitch
from .wan.model import NAGWanModelSwitch
# from .hunyuan_video.model import NAGHunyuanVideoSwitch
# from .hidream.model import NAGHiDreamImageTransformer2DModelSwitch
def sample_with_nag(
model,
noise,
positive, negative, nag_negative,
cfg,
nag_scale, nag_tau, nag_alpha, nag_sigma_end,
device,
sampler,
sigmas,
model_options={},
latent_image=None, denoise_mask=None, callback=None, disable_pbar=False, seed=None,
latent_shapes=None, **kwargs,
):
guider = NAGCFGGuider(model)
guider.set_conds(positive, negative)
guider.set_cfg(cfg)
guider.set_batch_size(latent_image.shape[0])
guider.set_nag(nag_negative, nag_scale, nag_tau, nag_alpha, nag_sigma_end)
return guider.sample(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes, **kwargs)
class NAGCFGGuider(CFGGuider):
def __init__(self, model_patcher: ModelPatcher):
super().__init__(model_patcher=model_patcher)
self.origin_nag_negative_cond = None
self.nag_scale = 5.0
self.nag_tau = 3.5
self.nag_alpha = 0.25
self.nag_sigma_end = 0.
self.batch_size = 1
def set_conds(self, positive, negative=None):
self.inner_set_conds(
{"positive": positive, "negative": negative} if negative is not None else {"positive": positive})
def set_batch_size(self, batch_size):
self.batch_size = batch_size
def set_nag(self, nag_negative_cond, nag_scale, nag_tau, nag_alpha, nag_sigma_end):
self.origin_nag_negative_cond = nag_negative_cond
self.nag_scale = nag_scale
self.nag_tau = nag_tau
self.nag_alpha = nag_alpha
self.nag_sigma_end = nag_sigma_end
def __call__(self, *args, **kwargs):
return self.predict_noise(*args, **kwargs)
def inner_sample(self, noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes=None, **kwargs):
if latent_image is not None and torch.count_nonzero(latent_image) > 0: #Don't shift the empty latent image.
latent_image = self.inner_model.process_latent_in(latent_image)
self.conds = process_conds(self.inner_model, noise, self.conds, device, latent_image, denoise_mask, seed)
extra_model_options = comfy.model_patcher.create_model_options_clone(self.model_options)
extra_model_options.setdefault("transformer_options", {})["sample_sigmas"] = sigmas
extra_args = {"model_options": extra_model_options, "seed": seed}
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
sampler.sample,
sampler,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, extra_args["model_options"], is_model_options=True)
)
samples = executor.execute(self, sigmas, extra_args, callback, noise, latent_image, denoise_mask, disable_pbar)
return self.inner_model.process_latent_out(samples.to(torch.float32))
def sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None, latent_shapes=None, **kwargs):
if sigmas.shape[-1] == 0:
return latent_image
self.conds = {}
for k in self.original_conds:
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
preprocess_conds_hooks(self.conds)
apply_guidance = self.nag_scale > 1.
self.nag_negative_cond = None
if apply_guidance:
self.nag_negative_cond = copy.deepcopy(self.origin_nag_negative_cond)
model = self.model_patcher.model.diffusion_model
if isinstance(model, OptimizedModule):
model = model._orig_mod
model_type = type(model)
# if model_type == Flux:
# switcher_cls = NAGFluxSwitch
# elif model_type == Chroma:
# switcher_cls = NAGChromaSwitch
# elif model_type == UNetModel:
# switcher_cls = NAGUNetModelSwitch
# elif model_type == OpenAISignatureMMDITWrapper:
# switcher_cls = NAGOpenAISignatureMMDITWrapperSwitch
if model_type in [WanModel, VaceWanModel]:
switcher_cls = NAGWanModelSwitch
# elif model_type == HunyuanVideo:
# switcher_cls = NAGHunyuanVideoSwitch
# elif model_type == HiDreamImageTransformer2DModel:
# switcher_cls = NAGHiDreamImageTransformer2DModelSwitch
else:
raise ValueError(
f"Model type {model_type} is not support for NAGCFGGuider"
)
self.nag_negative_cond[0][0] = self.nag_negative_cond[0][0].expand(self.batch_size, -1, -1)
if self.nag_negative_cond[0][1].get("pooled_output", None) is not None:
self.nag_negative_cond[0][1]["pooled_output"] = self.nag_negative_cond[0][1]["pooled_output"].expand(self.batch_size, -1)
switcher = switcher_cls(
model,
self.nag_negative_cond,
self.nag_scale, self.nag_tau, self.nag_alpha, self.nag_sigma_end,
)
switcher.set_nag()
try:
orig_model_options = self.model_options
self.model_options = comfy.model_patcher.create_model_options_clone(self.model_options)
# if one hook type (or just None), then don't bother caching weights for hooks (will never change after first step)
orig_hook_mode = self.model_patcher.hook_mode
if get_total_hook_groups_in_conds(self.conds) <= 1:
self.model_patcher.hook_mode = comfy.hooks.EnumHookMode.MinVram
comfy.sampler_helpers.prepare_model_patcher(self.model_patcher, self.conds, self.model_options)
filter_registered_hooks_on_conds(self.conds, self.model_options)
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
self.outer_sample,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, self.model_options, is_model_options=True)
)
output = executor.execute(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
finally:
cast_to_load_options(self.model_options, device=self.model_patcher.offload_device)
self.model_options = orig_model_options
self.model_patcher.hook_mode = orig_hook_mode
self.model_patcher.restore_hook_patches()
if apply_guidance:
switcher.set_origin()
del self.conds
del self.nag_negative_cond
return output
class KSamplerWithNAG(KSampler):
def sample(
self,
noise,
positive, negative, nag_negative,
cfg,
nag_scale, nag_tau, nag_alpha, nag_sigma_end,
latent_image=None,
start_step=None, last_step=None, force_full_denoise=False,
denoise_mask=None,
sigmas=None, callback=None, disable_pbar=False, seed=None,
latent_shapes=None,
**kwargs,
):
if sigmas is None:
sigmas = self.sigmas
if last_step is not None and last_step < (len(sigmas) - 1):
sigmas = sigmas[:last_step + 1]
if force_full_denoise:
sigmas[-1] = 0
if start_step is not None:
if start_step < (len(sigmas) - 1):
sigmas = sigmas[start_step:]
else:
if latent_image is not None:
return latent_image
else:
return torch.zeros_like(noise)
sampler = sampler_object(self.sampler)
return sample_with_nag(
self.model,
noise,
positive, negative, nag_negative,
cfg,
nag_scale, nag_tau, nag_alpha, nag_sigma_end,
self.device,
sampler,
sigmas,
self.model_options,
latent_image=latent_image, denoise_mask=denoise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed,
latent_shapes=latent_shapes, **kwargs,
)
+96
View File
@@ -0,0 +1,96 @@
import math
import torch
def nag(z_positive, z_negative, scale, tau, alpha):
z_guidance = z_positive * scale - z_negative * (scale - 1)
norm_positive = torch.norm(z_positive, p=1, dim=-1, keepdim=True).expand(*z_positive.shape)
norm_guidance = torch.norm(z_guidance, p=1, dim=-1, keepdim=True).expand(*z_guidance.shape)
scale = norm_guidance / norm_positive
z_guidance = z_guidance * torch.minimum(scale, scale.new_ones(1) * tau) / scale
z_guidance = z_guidance * alpha + z_positive * (1 - alpha)
return z_guidance
def cat_context(context, nag_negative_context, trim_context=False, dim=1):
assert dim in [1, 2]
nag_negative_context = nag_negative_context.to(context)
context_len = context.shape[dim]
nag_neg_context_len = nag_negative_context.shape[dim]
if context_len < nag_neg_context_len:
if dim == 1:
context = context.repeat(1, math.ceil(nag_neg_context_len / context_len), 1)
if trim_context:
context = context[:, -nag_neg_context_len:]
else:
context = context.repeat(1, 1, math.ceil(nag_neg_context_len / context_len), 1)
if trim_context:
context = context[:, :, -nag_neg_context_len:]
context_len = context.shape[dim]
if dim == 1:
nag_negative_context = nag_negative_context.repeat(1, math.ceil(context_len / nag_neg_context_len), 1)
nag_negative_context = nag_negative_context[:, -context_len:]
else:
nag_negative_context = nag_negative_context.repeat(1, 1, math.ceil(context_len / nag_neg_context_len), 1)
nag_negative_context = nag_negative_context[:, :, -context_len:]
return torch.cat([context, nag_negative_context], dim=0)
def check_nag_activation(transformer_options, nag_sigma_end):
apply_nag = torch.all(transformer_options["sigmas"] >= nag_sigma_end)
positive_batch = 0 in transformer_options["cond_or_uncond"]
return apply_nag and positive_batch
def get_closure_vars(func):
if func.__closure__ is None:
return {}
return {
var: cell.cell_contents
for var, cell in zip(func.__code__.co_freevars, func.__closure__)
}
def is_from_wavespeed(func):
closure = get_closure_vars(func)
return "residual_diff_threshold" in closure \
and "validate_can_use_cache_function" in closure
class NAGSwitch:
def __init__(
self,
model: torch.nn.Module,
nag_negative_cond,
nag_scale, nag_tau, nag_alpha, nag_sigma_end,
):
self.model = model
self.nag_negative_cond = nag_negative_cond
self.nag_scale = nag_scale
self.nag_tau = nag_tau
self.nag_alpha = nag_alpha
self.nag_sigma_end = nag_sigma_end
self.origin_forward = model.forward
def set_nag(self):
pass
def set_origin(self):
self.model.forward = self.origin_forward
# https://github.com/welltop-cn/ComfyUI-TeaCache/blob/4bca908bf53b029ea5739cb69ef2a9e6c06e6752/nodes.py
def poly1d(coefficients, x):
result = torch.zeros_like(x)
for i, coeff in enumerate(coefficients):
result += coeff * (x ** (len(coefficients) - 1 - i))
return result
+719
View File
@@ -0,0 +1,719 @@
from types import MethodType
from functools import partial
import torch
from einops import repeat
import comfy
from comfy.ldm.modules.attention import optimized_attention
from comfy.ldm.wan.model import (
WanModel,
VaceWanModel,
WanSelfAttention,
WanT2VCrossAttention,
WanI2VCrossAttention,
sinusoidal_embedding_1d,
)
from ..utils import nag, cat_context, check_nag_activation, poly1d, NAGSwitch
class NAGWanT2VCrossAttention(WanT2VCrossAttention):
def __init__(
self,
*args,
nag_scale: float = 1,
nag_tau: float = 3.5,
nag_alpha: float = 0.5,
**kwargs,
):
super().__init__(*args, **kwargs)
self.nag_scale = nag_scale
self.nag_tau = nag_tau
self.nag_alpha = nag_alpha
def forward(
self,
x,
context,
context_pad_len: int = None,
nag_pad_len: int = None,
**kwargs,
):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
"""
origin_bsz = len(context) - len(x)
assert origin_bsz != 0
# compute query, key, value
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(context))
v = self.v(context)
q_negative = q[-origin_bsz:]
k, k_negative = k[:-origin_bsz, :, context_pad_len:], k[-origin_bsz:, :, nag_pad_len:]
v, v_negative = v[:-origin_bsz, :, context_pad_len:], v[-origin_bsz:, :, nag_pad_len:]
# compute attention
x = optimized_attention(q, k, v, heads=self.num_heads)
x_negative = optimized_attention(q_negative, k_negative, v_negative, heads=self.num_heads)
x_positive = x[-origin_bsz:]
x_guidance = nag(x_positive, x_negative, self.nag_scale, self.nag_tau, self.nag_alpha)
x = torch.cat([x[:-origin_bsz], x_guidance], dim=0)
x = self.o(x)
return x
class NAGWanI2VCrossAttention(WanI2VCrossAttention):
def __init__(
self,
*args,
nag_scale: float = 1,
nag_tau: float = 3.5,
nag_alpha: float = 0.5,
**kwargs,
):
super().__init__(*args, **kwargs)
self.nag_scale = nag_scale
self.nag_tau = nag_tau
self.nag_alpha = nag_alpha
def forward(
self,
x,
context,
context_img_len,
context_pad_len: int = None,
nag_pad_len: int = None,
):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
"""
origin_bsz = len(context) - len(x)
assert origin_bsz != 0
context_img = context[:, :context_img_len]
context = context[:, context_img_len:]
# compute query, key, value
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(context))
v = self.v(context)
k_img = self.norm_k_img(self.k_img(context_img))
v_img = self.v_img(context_img)
q_negative = q[-origin_bsz:]
k, k_negative = k[:-origin_bsz, :, context_pad_len:], k[-origin_bsz:, :, nag_pad_len:]
v, v_negative = v[:-origin_bsz, :, context_pad_len:], v[-origin_bsz:, :, nag_pad_len:]
k_img, k_img_negative = k_img[:-origin_bsz], k_img[-origin_bsz:]
v_img, v_img_negative = v_img[:-origin_bsz], v_img[-origin_bsz:]
img_x = optimized_attention(q, k_img, v_img, heads=self.num_heads)
img_x_negative = optimized_attention(q_negative, k_img_negative, v_img_negative, heads=self.num_heads)
x = optimized_attention(q, k, v, heads=self.num_heads)
x_negative = optimized_attention(q_negative, k_negative, v_negative, heads=self.num_heads)
x_positive = x[-origin_bsz:]
x_guidance = nag(x_positive, x_negative, self.nag_scale, self.nag_tau, self.nag_alpha)
x = torch.cat([x[:-origin_bsz], x_guidance], dim=0)
img_x_positive = img_x[-origin_bsz:]
img_x_guidance = nag(img_x_positive, img_x_negative, self.nag_scale, self.nag_tau, self.nag_alpha)
img_x = torch.cat([img_x[:-origin_bsz], img_x_guidance], dim=0)
# output
x = x + img_x
x = self.o(x)
return x
class NAGWanModel(WanModel):
def forward_orig(
self,
x,
t,
context,
clip_fea=None,
freqs=None,
transformer_options={},
**kwargs,
):
r"""
Forward pass through the diffusion model
Args:
x (Tensor):
List of input video tensors with shape [B, C_in, F, H, W]
t (Tensor):
Diffusion timesteps tensor of shape [B]
context (List[Tensor]):
List of text embeddings each with shape [B, L, C]
seq_len (`int`):
Maximum sequence length for positional encoding
clip_fea (Tensor, *optional*):
CLIP image features for image-to-video mode
y (List[Tensor], *optional*):
Conditional video inputs for image-to-video mode, same shape as x
Returns:
List[Tensor]:
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
# time embeddings
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).to(dtype=x[0].dtype))
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
# context
context = self.text_embedding(context)
context_img_len = None
if clip_fea is not None:
if self.img_emb is not None:
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context_clip = torch.cat([context_clip, context_clip[-context.shape[0] - context_clip.shape[0]:]])
context = torch.concat([context_clip, context], dim=1)
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
blocks_replace = patches_replace.get("dit", {})
for i, block in enumerate(self.blocks):
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["img"] = block(args["img"], context=args["txt"], e=args["vec"], freqs=args["pe"],
context_img_len=context_img_len)
return out
out = blocks_replace[("double_block", i)]({"img": x, "txt": context, "vec": e0, "pe": freqs},
{"original_block": block_wrap})
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
# head
x = self.head(x, e)
# unpatchify
x = self.unpatchify(x, grid_sizes)
return x
def forward_orig_with_teacache(
self,
x,
t,
context,
clip_fea=None,
freqs=None,
transformer_options={},
**kwargs,
):
r"""
Forward pass through the diffusion model
Args:
x (Tensor):
List of input video tensors with shape [B, C_in, F, H, W]
t (Tensor):
Diffusion timesteps tensor of shape [B]
context (List[Tensor]):
List of text embeddings each with shape [B, L, C]
seq_len (`int`):
Maximum sequence length for positional encoding
clip_fea (Tensor, *optional*):
CLIP image features for image-to-video mode
y (List[Tensor], *optional*):
Conditional video inputs for image-to-video mode, same shape as x
Returns:
List[Tensor]:
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
enable_teacache = transformer_options.get("enable_teacache", True)
rel_l1_thresh = transformer_options.get("rel_l1_thresh")
coefficients = transformer_options.get("coefficients")
cond_or_uncond = transformer_options.get("cond_or_uncond")
model_type = transformer_options.get("model_type")
cache_device = transformer_options.get("cache_device")
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
# time embeddings
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).to(dtype=x[0].dtype))
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
# context
context = self.text_embedding(context)
context_img_len = None
if clip_fea is not None:
if self.img_emb is not None:
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context_clip = torch.cat([context_clip, context_clip[-context.shape[0] - context_clip.shape[0]:]])
context = torch.concat([context_clip, context], dim=1)
context_img_len = clip_fea.shape[-2]
patches_replace = transformer_options.get("patches_replace", {})
blocks_replace = patches_replace.get("dit", {})
if enable_teacache:
modulated_inp = e0.to(cache_device) if "ret_mode" in model_type else e.to(cache_device)
if not hasattr(self, 'teacache_state'):
self.teacache_state = {
0: {'should_calc': True, 'accumulated_rel_l1_distance': 0, 'previous_modulated_input': None,
'previous_residual': None},
1: {'should_calc': True, 'accumulated_rel_l1_distance': 0, 'previous_modulated_input': None,
'previous_residual': None}
}
def update_cache_state(cache, modulated_inp):
if cache['previous_modulated_input'] is not None:
try:
cache['accumulated_rel_l1_distance'] += poly1d(coefficients, (
(modulated_inp - cache['previous_modulated_input']).abs().mean() / cache[
'previous_modulated_input'].abs().mean()))
if cache['accumulated_rel_l1_distance'] < rel_l1_thresh:
cache['should_calc'] = False
else:
cache['should_calc'] = True
cache['accumulated_rel_l1_distance'] = 0
except:
cache['should_calc'] = True
cache['accumulated_rel_l1_distance'] = 0
cache['previous_modulated_input'] = modulated_inp
b = int(len(x) / len(cond_or_uncond))
for i, k in enumerate(cond_or_uncond):
update_cache_state(self.teacache_state[k], modulated_inp[i * b:(i + 1) * b])
if enable_teacache:
should_calc = False
for k in cond_or_uncond:
should_calc = (should_calc or self.teacache_state[k]['should_calc'])
else:
should_calc = True
if should_calc:
ori_x = x.to(cache_device)
else:
should_calc = False
if should_calc:
for i, block in enumerate(self.blocks):
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["img"] = block(args["img"], context=args["txt"], e=args["vec"], freqs=args["pe"],
context_img_len=context_img_len)
return out
out = blocks_replace[("double_block", i)]({"img": x, "txt": context, "vec": e0, "pe": freqs},
{"original_block": block_wrap})
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
else:
for i, k in enumerate(cond_or_uncond):
x[i * b:(i + 1) * b] += self.teacache_state[k]['previous_residual'].to(x.device)
if enable_teacache and should_calc:
for i, k in enumerate(cond_or_uncond):
self.teacache_state[k]['previous_residual'] = (x.to(cache_device) - ori_x)[i*b:(i+1)*b]
# head
x = self.head(x, e)
# unpatchify
x = self.unpatchify(x, grid_sizes)
return x
def forward(
self,
x,
timestep,
context,
clip_fea=None,
time_dim_concat=None,
transformer_options={},
nag_negative_context=None,
nag_sigma_end=0.,
**kwargs,
):
apply_nag = check_nag_activation(transformer_options, nag_sigma_end)
if apply_nag:
origin_context_len = context.shape[1]
context = cat_context(context, nag_negative_context, trim_context=True)
context_pad_len = context.shape[1] - origin_context_len
nag_pad_len = context.shape[1] - nag_negative_context.shape[1]
forward_orig_ = self.forward_orig
cross_attns_forward = list()
if transformer_options.get("enable_teacache", False):
self.forward_orig = MethodType(NAGWanModel.forward_orig_with_teacache, self)
else:
self.forward_orig = MethodType(NAGWanModel.forward_orig, self)
cross_attn_cls = NAGWanT2VCrossAttention if self.model_type == "t2v" else NAGWanI2VCrossAttention
for name, module in self.named_modules():
if "cross_attn" in name and isinstance(module, WanSelfAttention):
cross_attns_forward.append((module, module.forward))
module.forward = MethodType(
partial(
cross_attn_cls.forward,
context_pad_len=context_pad_len,
nag_pad_len=nag_pad_len,
),
module,
)
bs, c, t, h, w = x.shape
x = comfy.ldm.common_dit.pad_to_patch_size(x, self.patch_size)
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if time_dim_concat is not None:
time_dim_concat = comfy.ldm.common_dit.pad_to_patch_size(time_dim_concat, self.patch_size)
x = torch.cat([x, time_dim_concat], dim=2)
t_len = ((x.shape[2] + (patch_size[0] // 2)) // patch_size[0])
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, t_len - 1, steps=t_len, device=x.device,
dtype=x.dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device,
dtype=x.dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device,
dtype=x.dtype).reshape(1, 1, -1)
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=bs)
freqs = self.rope_embedder(img_ids).movedim(1, 2)
output = self.forward_orig(
x, timestep, context, clip_fea=clip_fea, freqs=freqs,
transformer_options=transformer_options, **kwargs)[:, :, :t, :h, :w]
if apply_nag:
self.forward_orig = forward_orig_
for mod, forward_fn in cross_attns_forward:
mod.forward = forward_fn
return output
class NAGVaceWanModel(VaceWanModel):
def forward_orig(
self,
x,
t,
context,
vace_context,
vace_strength,
clip_fea=None,
freqs=None,
transformer_options={},
**kwargs,
):
origin_batch_size = context.shape[0] - x.shape[0]
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
# time embeddings
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).to(dtype=x[0].dtype))
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
# context
context = self.text_embedding(context)
context_img_len = None
if clip_fea is not None:
if self.img_emb is not None:
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context_clip = torch.cat([context_clip, context_clip[-origin_batch_size:]])
context = torch.concat([context_clip, context], dim=1)
context_img_len = clip_fea.shape[-2]
orig_shape = list(vace_context.shape)
vace_context = vace_context.movedim(0, 1).reshape([-1] + orig_shape[2:])
c = self.vace_patch_embedding(vace_context.float()).to(vace_context.dtype)
c = c.flatten(2).transpose(1, 2)
c = list(c.split(orig_shape[0], dim=0))
# arguments
x_orig = x
patches_replace = transformer_options.get("patches_replace", {})
blocks_replace = patches_replace.get("dit", {})
for i, block in enumerate(self.blocks):
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["img"] = block(args["img"], context=args["txt"], e=args["vec"], freqs=args["pe"], context_img_len=context_img_len)
return out
out = blocks_replace[("double_block", i)]({"img": x, "txt": context, "vec": e0, "pe": freqs}, {"original_block": block_wrap})
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
ii = self.vace_layers_mapping.get(i, None)
if ii is not None:
for iii in range(len(c)):
c_skip, c[iii] = self.vace_blocks[ii](c[iii], x=x_orig, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
x += c_skip * vace_strength[iii]
del c_skip
# head
x = self.head(x, e)
# unpatchify
x = self.unpatchify(x, grid_sizes)
return x
def forward_orig_with_teacache(
self,
x,
t,
context,
vace_context,
vace_strength,
clip_fea=None,
freqs=None,
transformer_options={},
**kwargs,
):
enable_teacache = transformer_options.get("enable_teacache", True)
rel_l1_thresh = transformer_options.get("rel_l1_thresh")
coefficients = transformer_options.get("coefficients")
cond_or_uncond = transformer_options.get("cond_or_uncond")
model_type = transformer_options.get("model_type")
cache_device = transformer_options.get("cache_device")
origin_batch_size = context.shape[0] - x.shape[0]
# embeddings
x = self.patch_embedding(x.float()).to(x.dtype)
grid_sizes = x.shape[2:]
x = x.flatten(2).transpose(1, 2)
# time embeddings
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).to(dtype=x[0].dtype))
e0 = self.time_projection(e).unflatten(1, (6, self.dim))
# context
context = self.text_embedding(context)
context_img_len = None
if clip_fea is not None:
if self.img_emb is not None:
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context_clip = torch.cat([context_clip, context_clip[-origin_batch_size:]])
context = torch.concat([context_clip, context], dim=1)
context_img_len = clip_fea.shape[-2]
orig_shape = list(vace_context.shape)
vace_context = vace_context.movedim(0, 1).reshape([-1] + orig_shape[2:])
c = self.vace_patch_embedding(vace_context.float()).to(vace_context.dtype)
c = c.flatten(2).transpose(1, 2)
c = list(c.split(orig_shape[0], dim=0))
# arguments
x_orig = x
patches_replace = transformer_options.get("patches_replace", {})
blocks_replace = patches_replace.get("dit", {})
if enable_teacache:
modulated_inp = e0.to(cache_device) if "ret_mode" in model_type else e.to(cache_device)
if not hasattr(self, 'teacache_state'):
self.teacache_state = {
0: {'should_calc': True, 'accumulated_rel_l1_distance': 0, 'previous_modulated_input': None,
'previous_residual': None},
1: {'should_calc': True, 'accumulated_rel_l1_distance': 0, 'previous_modulated_input': None,
'previous_residual': None}
}
def update_cache_state(cache, modulated_inp):
if cache['previous_modulated_input'] is not None:
try:
cache['accumulated_rel_l1_distance'] += poly1d(coefficients, (
(modulated_inp - cache['previous_modulated_input']).abs().mean() / cache[
'previous_modulated_input'].abs().mean()))
if cache['accumulated_rel_l1_distance'] < rel_l1_thresh:
cache['should_calc'] = False
else:
cache['should_calc'] = True
cache['accumulated_rel_l1_distance'] = 0
except:
cache['should_calc'] = True
cache['accumulated_rel_l1_distance'] = 0
cache['previous_modulated_input'] = modulated_inp
b = int(len(x) / len(cond_or_uncond))
for i, k in enumerate(cond_or_uncond):
update_cache_state(self.teacache_state[k], modulated_inp[i * b:(i + 1) * b])
if enable_teacache:
should_calc = False
for k in cond_or_uncond:
should_calc = (should_calc or self.teacache_state[k]['should_calc'])
else:
should_calc = True
if should_calc:
ori_x = x.to(cache_device)
else:
should_calc = False
if should_calc:
for i, block in enumerate(self.blocks):
if ("double_block", i) in blocks_replace:
def block_wrap(args):
out = {}
out["img"] = block(args["img"], context=args["txt"], e=args["vec"], freqs=args["pe"], context_img_len=context_img_len)
return out
out = blocks_replace[("double_block", i)]({"img": x, "txt": context, "vec": e0, "pe": freqs}, {"original_block": block_wrap})
x = out["img"]
else:
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
ii = self.vace_layers_mapping.get(i, None)
if ii is not None:
for iii in range(len(c)):
c_skip, c[iii] = self.vace_blocks[ii](c[iii], x=x_orig, e=e0, freqs=freqs, context=context, context_img_len=context_img_len)
x += c_skip * vace_strength[iii]
del c_skip
else:
for i, k in enumerate(cond_or_uncond):
x[i * b:(i + 1) * b] += self.teacache_state[k]['previous_residual'].to(x.device)
if enable_teacache and should_calc:
for i, k in enumerate(cond_or_uncond):
self.teacache_state[k]['previous_residual'] = (x.to(cache_device) - ori_x)[i * b:(i + 1) * b]
# head
x = self.head(x, e)
# unpatchify
x = self.unpatchify(x, grid_sizes)
return x
def forward(
self,
x,
timestep,
context,
clip_fea=None,
time_dim_concat=None,
transformer_options={},
nag_negative_context=None,
nag_sigma_end=0.,
**kwargs,
):
apply_nag = check_nag_activation(transformer_options, nag_sigma_end)
if apply_nag:
origin_context_len = context.shape[1]
context = cat_context(context, nag_negative_context, trim_context=True)
context_pad_len = context.shape[1] - origin_context_len
nag_pad_len = context.shape[1] - nag_negative_context.shape[1]
forward_orig_ = self.forward_orig
cross_attns_forward = list()
if transformer_options.get("enable_teacache", False):
self.forward_orig = MethodType(NAGVaceWanModel.forward_orig_with_teacache, self)
else:
self.forward_orig = MethodType(NAGVaceWanModel.forward_orig, self)
for name, module in self.named_modules():
if "cross_attn" in name and isinstance(module, WanSelfAttention):
cross_attns_forward.append((module, module.forward))
module.forward = MethodType(
partial(
NAGWanT2VCrossAttention.forward,
context_pad_len=context_pad_len,
nag_pad_len=nag_pad_len,
),
module,
)
bs, c, t, h, w = x.shape
x = comfy.ldm.common_dit.pad_to_patch_size(x, self.patch_size)
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if time_dim_concat is not None:
time_dim_concat = comfy.ldm.common_dit.pad_to_patch_size(time_dim_concat, self.patch_size)
x = torch.cat([x, time_dim_concat], dim=2)
t_len = ((x.shape[2] + (patch_size[0] // 2)) // patch_size[0])
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, t_len - 1, steps=t_len, device=x.device,
dtype=x.dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device,
dtype=x.dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device,
dtype=x.dtype).reshape(1, 1, -1)
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=bs)
freqs = self.rope_embedder(img_ids).movedim(1, 2)
output = self.forward_orig(x, timestep, context, clip_fea=clip_fea, freqs=freqs,
transformer_options=transformer_options, **kwargs)[:, :, :t, :h, :w]
if apply_nag:
self.forward_orig = forward_orig_
for mod, forward_fn in cross_attns_forward:
mod.forward = forward_fn
return output
class NAGWanModelSwitch(NAGSwitch):
def set_nag(self):
nag_model_cls = NAGVaceWanModel if isinstance(self.model, VaceWanModel) else NAGWanModel
self.model.forward = MethodType(
partial(
nag_model_cls.forward,
nag_negative_context=self.nag_negative_cond[0][0],
nag_sigma_end=self.nag_sigma_end,
),
self.model,
)
for name, module in self.model.named_modules():
if "cross_attn" in name and isinstance(module, WanSelfAttention):
module.nag_scale = self.nag_scale
module.nag_tau = self.nag_tau
module.nag_alpha = self.nag_alpha
+853 -111
View File
File diff suppressed because it is too large Load Diff
+4 -2
View File
@@ -1,8 +1,8 @@
# pyproject.toml
[project]
name = "ComfyUI-SuperUltimateVaceTools"
description = "powerful nodes for wan2.1 vace"
version = "0.0.2"
description = "powerful nodes for wan2.1/wan2.2 vace"
version = "0.4.0"
license = { file = "LICENSE.txt" }
dependencies = []
@@ -13,3 +13,5 @@ Repository = "https://github.com/bbaudio-2025/ComfyUI-SuperUltimateVaceTools"
PublisherId = "bbaudio"
DisplayName = "ComfyUI-SuperUltimateVaceTools"
Icon = ""
+1
View File
@@ -0,0 +1 @@
color-matcher
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff