Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
93298fd928 | ||
|
|
74ed730b91 | ||
|
|
ca38b2868f | ||
|
|
054ab62e9f | ||
|
|
3a9686a010 | ||
|
|
52391ee62a | ||
|
|
58700c8099 | ||
|
|
07ed89b3cb | ||
|
|
b8b70c2b5c | ||
|
|
adb48b4e9f | ||
|
|
5679597026 | ||
|
|
3ca3c6d2c0 | ||
|
|
3bc9abe234 | ||
|
|
0f00c49f72 | ||
|
|
8801c9b4ea | ||
|
|
d76cd98c98 | ||
|
|
99820833b5 | ||
|
|
0fc8c5fe14 | ||
|
|
245a5920da | ||
|
|
d6e7c88881 | ||
|
|
3a5678505b | ||
|
|
8307159728 | ||
|
|
172e22d00a | ||
|
|
670875de24 | ||
|
|
fdd2e72d08 | ||
|
|
762bcb2013 | ||
|
|
ceb84ee3b2 | ||
|
|
8067f227c4 | ||
|
|
935ed652a0 | ||
|
|
38191b6f5e | ||
|
|
0f60de4d4b | ||
|
|
4991809da2 | ||
|
|
8ba542d03b | ||
|
|
22bea2b5f8 | ||
|
|
93900e1cf6 | ||
|
|
f999edee89 | ||
|
|
503e580a85 | ||
|
|
60facc895a | ||
|
|
923fa30e30 | ||
|
|
798ea62e44 | ||
|
|
27e1fb25f3 | ||
|
|
48faf262ec | ||
|
|
dd649aa05c | ||
|
|
a07081dccc | ||
|
|
719bdedfd6 | ||
|
|
71383e398c | ||
|
|
2ede0e1f3c | ||
|
|
5a875855a6 | ||
|
|
f5e7cae1ba | ||
|
|
9a033d9549 |
@@ -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`
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
+4
-2
@@ -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 = ""
|
||||
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user