Merge branch 'kijai:main' into main

This commit is contained in:
小六妞儿
2026-01-15 14:05:52 +08:00
committed by GitHub
14 changed files with 626 additions and 116 deletions
+199
View File
@@ -0,0 +1,199 @@
import torch
import torch.nn as nn
from einops import rearrange
from ..wanvideo.modules.attention import attention
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
return (x * (1 + scale) + shift)
def sinusoidal_embedding_1d(dim, position):
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x.to(position.dtype)
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
# 3d rope precompute
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
return f_freqs_cis, h_freqs_cis, w_freqs_cis
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
# 1d rope precompute
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
[: (dim // 2)].double() / dim))
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
return freqs_cis
def rope_apply(x, freqs, num_heads):
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
x.shape[0], x.shape[1], x.shape[2], -1, 2))
x_out = torch.view_as_real(x_out * freqs).flatten(2)
return x_out.to(x.dtype)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
def forward(self, x):
dtype = x.dtype
return self.norm(x.float()).to(dtype) * self.weight
class AttentionModule(nn.Module):
def __init__(self, num_heads, head_dim):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
def forward(self, q, k, v):
b, n, d = q.size(0), self.num_heads, self.head_dim
x = attention(
q.view(b, -1, n, d),
k.view(b, -1, n, d),
v.view(b, -1, n, d)
)
return x.flatten(2)
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x, freqs):
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(x))
v = self.v(x)
q = rope_apply(q, freqs, self.num_heads)
k = rope_apply(k, freqs, self.num_heads)
x = self.attn(q, k, v)
return self.o(x)
class CrossAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.k_img = nn.Linear(dim, dim)
self.v_img = nn.Linear(dim, dim)
self.norm_k_img = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
ctx = y
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(ctx))
v = self.v(ctx)
x = self.attn(q, k, v)
if clip_fea is not None:
k_img = self.norm_k_img(self.k_img(clip_fea))
v_img = self.v_img(clip_fea)
y = self.attn(q, k_img, v_img)
x = x + y
return self.o(x)
class GateModule(nn.Module):
def __init__(self,):
super().__init__()
def forward(self, x, gate, residual):
return x + gate * residual
class DiTBlock(nn.Module):
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.ffn_dim = ffn_dim
self.self_attn = SelfAttention(dim, num_heads, eps)
self.cross_attn = CrossAttention(dim, num_heads, eps)
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm3 = nn.LayerNorm(dim, eps=eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
approximate='tanh'), nn.Linear(ffn_dim, dim))
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.gate = GateModule()
def forward(self, x, context, t_mod, freqs, clip_fea=None):
has_seq = len(t_mod.shape) == 4
chunk_dim = 2 if has_seq else 1
# msa: multi-head self-attention mlp: multi-layer perceptron
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
if has_seq:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
)
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
x = self.gate(x, gate_mlp, self.ffn(input_x))
return x
class WanModelDualControl(torch.nn.Module):
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
super().__init__()
self.control_layers = control_layers
self.control_blocks_dense = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_blocks_sparse = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
self.control_text_linear = torch.nn.Linear(dim, dim//2)
self.control_t_mod = torch.nn.Linear(dim, dim//2)
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
head_dim = dim // num_heads
self.freqs = precompute_freqs_cis_3d(head_dim)
+88
View File
@@ -0,0 +1,88 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddDualControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
},
"optional": {
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
updated = dict(embeds)
updated.setdefault("dual_control", {})
if dense is None and sparse is None:
raise ValueError("At least one of dense or sparse inputs must be provided.")
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
height = dense.shape[1] if dense is not None else sparse.shape[1]
width = dense.shape[2] if dense is not None else sparse.shape[2]
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
msk = msk.transpose(1, 2)
dense_input_latent = sparse_input_latent = None
vae.to(device)
if dense is not None:
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
if sparse is not None:
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
if prev_images is not None:
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
vae.to(offload_device)
updated["dual_control"]["dense_input_latent"] = dense_input_latent
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
updated["dual_control"]["strength"] = strength
updated["dual_control"]["start_percent"] = start_percent
updated["dual_control"]["end_percent"] = end_percent
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
}
+5 -4
View File
@@ -1,11 +1,11 @@
try:
from .utils import check_duplicate_nodes, log
from .utils import check_duplicate_nodes, log, color_text
duplicate_dirs = check_duplicate_nodes()
if duplicate_dirs:
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
for dir_path in duplicate_dirs:
warning_msg += f" - {dir_path}\n"
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except:
pass
@@ -49,6 +49,7 @@ OPTIONAL_MODULES = [
(".WanMove.nodes", "WanMove"),
(".SCAIL.nodes", "SCAIL"),
(".LongCat.nodes", "LongCat"),
(".LongVie2.nodes", "LongVie2"),
]
def register_nodes(module_path: str, name: str, optional: bool) -> None:
@@ -71,4 +72,4 @@ for module_path, name in REQUIRED_MODULES:
for module_path, name in OPTIONAL_MODULES:
register_nodes(module_path, name, optional=True)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+1 -1
View File
@@ -56,7 +56,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert:
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
weight_key = module_prefix + "weight"
if weight_key not in state_dict:
continue
+1 -1
View File
@@ -16,7 +16,7 @@ import copy
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
vae_upscale_factor = 16
vae_upscale_factor = 8
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
+101 -20
View File
@@ -13,7 +13,7 @@ from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
from .custom_linear import _replace_linear
from accelerate import init_empty_weights
from .utils import set_module_tensor_to_device
from .utils import set_module_tensor_to_device, get_module_memory_mb_per_device
import folder_paths
import comfy.model_management as mm
@@ -36,6 +36,9 @@ try:
except:
PromptServer = None
attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled",
"sageattn_ultravico", "comfy"]
#from city96's gguf nodes
def update_folder_names_and_paths(key, targets=[]):
# check for existing key
@@ -827,7 +830,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
all_tensors.extend(r.tensors)
for tensor in all_tensors:
name = rename_fuser_block(tensor.name)
if "glob" not in name and "audio_proj" in name:
if "glob" not in name and "multitalk_audio_proj" not in name and "audio_proj" in name:
name = name.replace("audio_proj", "multitalk_audio_proj")
load_device = device
if "vace_blocks." in name:
@@ -861,7 +864,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
)
transformer.gguf_patched = True
else:
log.info("Using accelerate to load and assign model weights to device...")
log.info("Loading and assigning model weights to device...")
named_params = transformer.named_parameters()
for name, param in tqdm(named_params,
@@ -922,6 +925,11 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
pbar.update(100)
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
memory_on_device = get_module_memory_mb_per_device(transformer)
log.info("-" * 25)
log.info("Transformer weights loaded:")
for dev, mem_mb in memory_on_device.items():
log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB")
pbar.update_absolute(0)
@@ -1006,6 +1014,66 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
del lora_sd
return patcher, control_lora, unianimate_sd
class WanVideoSetAttentionModeOverride:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL", ),
"attention_mode": (attention_modes, {"default": "sdpa"}),
"start_step": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Step to start applying the attention mode override"}),
"end_step": ("INT", {"default": 10000, "min": 1, "max": 10000, "step": 1, "tooltip": "Step to end applying the attention mode override"}),
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Print verbose info about attention mode override during generation"}),
},
"optional": {
"blocks":("INT", {"forceInput": True} ),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "getmodelpath"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Override the attention mode for the model for specific step and/or block range"
def getmodelpath(self, model, attention_mode, start_step, end_step, verbose, blocks=None):
model_clone = model.clone()
attention_mode_override = {
"mode": attention_mode,
"start_step": start_step,
"end_step": end_step,
"verbose": verbose,
}
if blocks is not None:
attention_mode_override["blocks"] = blocks
model_clone.model_options['transformer_options']["attention_mode_override"] = attention_mode_override
return (model_clone,)
class WanVideoUltraVicoSettings:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL", ),
"alpha": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Alpha value for the decay, higher values mean slower decay"}),
},
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "getmodelpath"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Set UltraVico parameters, attention mode still needs to be set to sageattn_ultravico, https://github.com/thu-ml/DiT-Extrapolation"
def getmodelpath(self, model, alpha):
model_clone = model.clone()
model_clone.model_options['transformer_options']["ultravico_alpha"] = alpha
return (model_clone,)
#region Model loading
class WanVideoModelLoader:
@classmethod
@@ -1020,17 +1088,7 @@ class WanVideoModelLoader:
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
},
"optional": {
"attention_mode": ([
"sdpa",
"flash_attn_2",
"flash_attn_3",
"sageattn",
"sageattn_3",
"radial_sage_attention",
"sageattn_compiled",
"sageattn_ultravico",
"comfy"
], {"default": "sdpa"}),
"attention_mode": (attention_modes, {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
"lora": ("WANVIDLORA", {"default": None}),
@@ -1235,9 +1293,7 @@ class WanVideoModelLoader:
lynx_ip_layers = "lite"
model_type = "t2v"
if "audio_injector.injector.0.k.weight" in sd:
model_type = "s2v"
elif not "text_embedding.0.weight" in sd:
if not "text_embedding.0.weight" in sd:
model_type = "no_cross_attn" #minimaxremover
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
model_type = "fl2v"
@@ -1247,6 +1303,8 @@ class WanVideoModelLoader:
model_type = "t2v"
elif "control_adapter.conv.weight" in sd:
model_type = "t2v"
if "audio_injector.injector.0.k.weight" in sd:
model_type = "s2v"
out_dim = 16
if dim == 5120: #14B
@@ -1623,6 +1681,24 @@ class WanVideoModelLoader:
block.ref_attn_v_img = nn.Linear(in_features, out_features)
block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6)
if "blocks.0.control_blocks_dense.cross_attn.k.weight" in sd:
log.info("LongVie2 model detected, patching model...")
from .LongVie2.modules import WanModelDualControl
control_layers = 12
with init_empty_weights():
dual_controller = WanModelDualControl(dim=5120, ffn_dim=13824, eps=1e-06, num_heads=40, control_layers=control_layers)
for b in range(control_layers):
transformer.blocks[b].control_blocks_dense = dual_controller.control_blocks_dense[b]
transformer.blocks[b].control_blocks_sparse = dual_controller.control_blocks_sparse[b]
transformer.blocks[b].control_combine_linears = dual_controller.control_combine_linears[b]
transformer.dual_controller = nn.Module()
transformer.dual_controller.control_initial_combine_linear_dense = dual_controller.control_initial_combine_linear_dense
transformer.dual_controller.control_initial_combine_linear_sparse = dual_controller.control_initial_combine_linear_sparse
transformer.dual_controller.control_t_mod = dual_controller.control_t_mod
transformer.dual_controller.control_text_linear = dual_controller.control_text_linear
transformer.dual_controller_freqs = dual_controller.freqs
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
@@ -1803,6 +1879,7 @@ class WanVideoVAELoader:
),
"compile_args": ("WANCOMPILEARGS", ),
"use_cpu_cache": ("BOOLEAN", {"default": False, "tooltip": "Reduces VRAM usage, but slows the VAE down a lot"}),
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Enables memory usage logging when using the model"}),
}
}
@@ -1812,7 +1889,7 @@ class WanVideoVAELoader:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False):
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False, verbose=False):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("vae", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
@@ -1829,9 +1906,9 @@ class WanVideoVAELoader:
pruning_rate = 0.0
if vae_sd["model.conv2.weight"].shape[0] == 16:
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
elif vae_sd["model.conv2.weight"].shape[0] == 48:
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
vae.load_state_dict(vae_sd)
del vae_sd
@@ -2043,6 +2120,8 @@ NODE_CLASS_MAPPINGS = {
"WanVideoTorchCompileSettings": WanVideoTorchCompileSettings,
"LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder,
"LoadWanVideoClipTextEncoder": LoadWanVideoClipTextEncoder,
"WanVideoSetAttentionModeOverride": WanVideoSetAttentionModeOverride,
"WanVideoUltraVicoSettings": WanVideoUltraVicoSettings,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2061,4 +2140,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoTorchCompileSettings": "WanVideo Torch Compile Settings",
"LoadWanVideoT5TextEncoder": "WanVideo T5 Text Encoder Loader",
"LoadWanVideoClipTextEncoder": "WanVideo CLIP Text Encoder Loader",
"WanVideoSetAttentionModeOverride": "WanVideo Set Attention Mode Override",
"WanVideoUltraVicoSettings": "WanVideo UltraVico Settings"
}
+32 -6
View File
@@ -96,7 +96,7 @@ class WanVideoSampler:
vae = image_embeds.get("vae", None)
tiled_vae = image_embeds.get("tiled_vae", False)
transformer_options = patcher.model_options.get("transformer_options", None)
transformer_options = copy.deepcopy(patcher.model_options.get("transformer_options", None))
merge_loras = transformer_options["merge_loras"]
block_swap_args = transformer_options.get("block_swap_args", None)
@@ -177,7 +177,6 @@ class WanVideoSampler:
start_step = scheduler.get("start_step", start_step)
elif scheduler != "multitalk":
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas, log_timesteps=True)
log.info(f"sigmas: {sample_scheduler.sigmas}")
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
@@ -241,8 +240,6 @@ class WanVideoSampler:
else:
image_cond[:, 1:] = 0
log.info(f"image_cond shape: {image_cond.shape}")
#ATI tracks
if transformer_options is not None:
ATI_tracks = transformer_options.get("ati_tracks", None)
@@ -1151,6 +1148,18 @@ class WanVideoSampler:
if context_options is None:
image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
# LongVie2 dual control
dual_control_embeds = image_embeds.get("dual_control", None)
if dual_control_embeds is not None and context_options is None:
dual_control_input = dict_to_device(dual_control_embeds.copy(), device, dtype) if dual_control_embeds is not None else None
prev_latents = dual_control_input.get("prev_latent", None)
if prev_latents is not None:
_sigma = dual_control_embeds.get("first_frame_noise_level", 0.925926)
log.info(f"Using dual control previous latents with first frame noise level: {_sigma}")
latent[:, :1] = (1 - _sigma) * prev_latents[:, -1:].to(latent) + _sigma * noise[:, :1]
prev_ones = torch.ones(20, *prev_latents.shape[1:], device=device, dtype=dtype)
dual_control_input["prev_latent"] = torch.cat([prev_ones, prev_latents]).unsqueeze(0)
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1403,6 +1412,19 @@ class WanVideoSampler:
if wanmove_embeds is not None and context_window is not None:
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
dual_control_in = None
if dual_control_embeds is not None:
if context_window is not None:
dual_control_in = dual_control_embeds.copy()
dense_input_latent = dual_control_embeds.get("dense_input_latent", None)
if dense_input_latent is not None:
dual_control_in["dense_input_latent"] = dual_control_embeds["dense_input_latent"][:, :, context_window]
sparse_input_latent = dual_control_embeds.get("sparse_input_latent", None)
if sparse_input_latent is not None:
dual_control_in["sparse_input_latent"] = dual_control_embeds["sparse_input_latent"][:, :, context_window]
else:
dual_control_in = dual_control_input
base_params = {
'x': [z], # latent
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
@@ -1465,6 +1487,8 @@ class WanVideoSampler:
"one_to_all_input": one_to_all_data, # One-to-All input
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
"scail_input": scail_data_in, # SCAIL input
"dual_control_input": dual_control_in, # LongVie2 dual control input
"transformer_options": transformer_options
}
batch_size = 1
@@ -1678,8 +1702,8 @@ class WanVideoSampler:
callback = prepare_callback(patcher, len(timesteps))
if not multitalk_sampling and not framepack and not wananimate_loop:
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps-ttm_start_step} steps")
log.info("-" * 10 + " Sampling start " + "-" * 10)
log.info(f"{(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} (Input sequence length: {seq_len}) with {steps-ttm_start_step} steps")
# Differential diffusion prep
@@ -2579,6 +2603,8 @@ class WanVideoSampler:
if story_mem_latents is not None:
latent = latent[:, story_mem_latents.shape[1]:]
log.info("-" * 10 + " Sampling end " + "-" * 12)
cache_states = None
if cache_args is not None:
cache_report(transformer, cache_args)
+2 -2
View File
@@ -38,7 +38,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
window_th = 1560 * 21 / 2
window_th = frame_tokens * window_width / 2
dist2 = tl.abs(m - n).to(tl.int32)
dist_mask = dist2 <= window_th
@@ -46,7 +46,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
window3 = (m <= frame_tokens) & (n > 21*frame_tokens)
window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
qk = tl.where(window3, -1e4, qk)
+32 -5
View File
@@ -25,6 +25,23 @@ try:
except:
pass
COLOR_CODES = {
"reset": "\033[0m",
"red": "\033[31m",
"green": "\033[32m",
"yellow": "\033[33m",
"blue": "\033[34m",
"magenta": "\033[35m",
"cyan": "\033[36m",
"white": "\033[37m",
}
def color_text(text, color):
try:
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
except Exception:
return text
class MetaParameter(torch.nn.Parameter):
def __new__(cls, dtype, quant_type=None):
data = torch.empty(0, dtype=dtype)
@@ -191,10 +208,8 @@ def check_diffusers_version():
raise AssertionError("diffusers is not installed.")
def print_memory(device, process="Sampling"):
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"[{process}] Allocated memory: {memory=:.3f} GB")
log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
@@ -207,6 +222,18 @@ def get_module_memory_mb(module):
memory += param.nelement() * param.element_size()
return memory / (1024 * 1024) # Convert to MB
def get_module_memory_mb_per_device(module):
memory_per_device = {}
memory = 0
for param in module.parameters():
if param.data is not None:
device = str(param.device)
memory += param.nelement() * param.element_size()
memory_per_device[device] = memory_per_device.get(device, 0) + memory
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
return memory_per_device
def get_tensor_memory(tensor):
memory_bytes = tensor.element_size() * tensor.nelement()
return f"{memory_bytes / (1024 * 1024):.2f} MB"
@@ -666,9 +693,9 @@ 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
@@ -676,7 +703,7 @@ def check_duplicate_nodes():
'wanvideo' in path.name.lower() and
'wrapper' in path.name.lower()):
wanvideo_dirs.append(str(path))
return wanvideo_dirs
#https://github.com/temporalscorerescaling/TSR/
+4 -4
View File
@@ -80,9 +80,9 @@ except:
try:
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
) -> torch.Tensor:
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor)
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
@sageattn_func_ultravico.register_fake
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
@@ -94,7 +94,7 @@ except:
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
attention_mode='sdpa', attn_mask=None, multi_factor=0.9, heads=128):
attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128):
if "flash" in attention_mode:
return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
@@ -108,7 +108,7 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k
elif attention_mode == 'sageattn':
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn_ultravico':
return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor).contiguous()
return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous()
elif attention_mode == 'comfy':
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
else: # sdpa
+114 -33
View File
@@ -467,7 +467,7 @@ class WanSelfAttention(nn.Module):
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0):
def forward(self, q, k, v, seq_lens, transformer_options={}, attention_mode_override=None, lynx_ref_feature=None, lynx_ref_scale=1.0, onetoall_ref=None, onetoall_ref_scale=1.0, frame_tokens=1536):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -482,7 +482,7 @@ class WanSelfAttention(nn.Module):
if self.ref_adapter is not None and lynx_ref_feature is not None:
ref_x = self.ref_adapter(self, q, lynx_ref_feature)
x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads)
x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads, frame_tokens=frame_tokens, transformer_options=transformer_options)
if self.ref_adapter is not None and lynx_ref_feature is not None:
x = x.add(ref_x, alpha=lynx_ref_scale)
@@ -497,7 +497,7 @@ class WanSelfAttention(nn.Module):
attention_mode = self.attention_mode
if attention_mode_override is not None:
attention_mode = attention_mode_override
# Concatenate main and IP keys/values for main attention
full_k = torch.cat([k, k_ip], dim=1)
full_v = torch.cat([v, v_ip], dim=1)
@@ -1006,6 +1006,7 @@ class WanAttentionBlock(nn.Module):
longcat_num_cond_latents=0, longcat_avatar_options=None, #longcat image cond amount
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
e_tr=None, tr_num=0, tr_start=0, #token replacement
attention_mode_override=None, frame_tokens=None, transformer_options={}
):
r"""
Args:
@@ -1150,6 +1151,10 @@ class WanAttentionBlock(nn.Module):
if enhance_enabled:
feta_scores = get_feta_scores(q, k)
if self.attention_mode == "sageattn_3" and attention_mode_override is None:
if current_step != 0 and not last_step:
attention_mode_override = "sageattn"
#self-attention
split_attn = (context is not None
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
@@ -1161,19 +1166,14 @@ class WanAttentionBlock(nn.Module):
y = self.self_attn.forward_split(q, k, v, seq_lens, grid_sizes, seq_chunks)
elif ref_target_masks is not None: #multi/infinite talk
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
elif self.attention_mode == "radial_sage_attention":
elif self.attention_mode == "radial_sage_attention" or attention_mode_override is not None and attention_mode_override == "radial_sage_attention":
if self.dense_block or self.dense_timesteps is not None and current_step < self.dense_timesteps:
if self.dense_attention_mode == "sparse_sage_attn":
y = self.self_attn.forward_radial(q, k, v, dense_step=True)
else:
y = self.self_attn.forward(q, k, v, seq_lens)
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override=attention_mode_override)
else:
y = self.self_attn.forward_radial(q, k, v, dense_step=False)
elif self.attention_mode == "sageattn_3":
if current_step != 0 and not last_step:
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn_3")
else:
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn")
elif x_ip is not None and self.kv_cache is None: #stand-in
# First pass: cache IP keys/values and compute attention
self.kv_cache = {"k_ip": k_ip.detach(), "v_ip": v_ip.detach()}
@@ -1184,18 +1184,18 @@ class WanAttentionBlock(nn.Module):
v_ip = self.kv_cache["v_ip"]
full_k = torch.cat([k, k_ip], dim=1)
full_v = torch.cat([v, v_ip], dim=1)
y = self.self_attn.forward(q, full_k, full_v, seq_lens)
y = self.self_attn.forward(q, full_k, full_v, seq_lens, attention_mode_override=attention_mode_override)
elif is_longcat and longcat_num_cond_latents > 0:
if longcat_num_cond_latents == 1:
num_cond_latents_thw = longcat_num_cond_latents * (N // num_latent_frames)
# process the noise tokens
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
# process the condition tokens
x_cond = self.self_attn.forward(
q[:, :num_cond_latents_thw].contiguous(),
k[:, :num_cond_latents_thw].contiguous(),
v[:, :num_cond_latents_thw].contiguous(),
seq_lens)
seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif longcat_num_cond_latents > 1: # video continuation
@@ -1224,12 +1224,12 @@ class WanAttentionBlock(nn.Module):
k_non_ref = k[:, num_ref_latents_thw:].contiguous()
v_non_ref = v[:, num_ref_latents_thw:].contiguous()
x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens) # q_front has attention with ref + cond + noisy
x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens) # q_back has attention with ref + cond + noisy
x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens) # q_mask has attention with cond+noisy
x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_front has attention with ref + cond + noisy
x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_back has attention with ref + cond + noisy
x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_mask has attention with cond+noisy
x_noise = torch.cat([x_noise_front, x_noise_maskref, x_noise_back], dim=1).contiguous()
else:
x_noise = self.self_attn.forward(q_noise, k, v, seq_lens)
x_noise = self.self_attn.forward(q_noise, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
# process the condition tokens
q_ref = q[:, :num_ref_latents_thw].contiguous()
k_ref = k[:, :num_ref_latents_thw].contiguous()
@@ -1237,13 +1237,14 @@ class WanAttentionBlock(nn.Module):
q_cond = q[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
k_cond = k[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
v_cond = v[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens)
x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens)
x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
# merge x_cond and x_noise
y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous()
else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale)
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale,
onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override, transformer_options=transformer_options, frame_tokens=frame_tokens)
del q, k, v
@@ -2018,24 +2019,24 @@ class WanModel(torch.nn.Module):
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
# Clamp blocks_to_swap to valid range
blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks)))
log.info(f"Swapping {blocks_to_swap} transformer blocks")
self.blocks_to_swap = blocks_to_swap
self.prefetch_blocks = prefetch_blocks
self.block_swap_debug = block_swap_debug
self.offload_img_emb = offload_img_emb
self.offload_txt_emb = offload_txt_emb
total_offload_memory = 0
total_main_memory = 0
# Calculate the index where swapping starts
swap_start_idx = len(self.blocks) - blocks_to_swap
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
block_memory = get_module_memory_mb(block)
if b < swap_start_idx:
block.to(self.main_device)
total_main_memory += block_memory
@@ -2050,13 +2051,13 @@ class WanModel(torch.nn.Module):
# Clamp vace_blocks_to_swap to valid range
vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks)))
self.vace_blocks_to_swap = vace_blocks_to_swap
# Calculate the index where VACE swapping starts
vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap
for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"):
block_memory = get_module_memory_mb(block)
if b < vace_swap_start_idx:
block.to(self.main_device)
total_main_memory += block_memory
@@ -2067,13 +2068,13 @@ class WanModel(torch.nn.Module):
mm.soft_empty_cache()
gc.collect()
log.info("----------------------")
log.info(f"Block swap memory summary:")
log.info("-" * 25)
log.info("Block swap memory summary:")
log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB")
log.info(f"Transformer blocks on {self.main_device}: {total_main_memory:.2f}MB")
log.info(f"Total memory used by transformer blocks: {(total_offload_memory + total_main_memory):.2f}MB")
log.info(f"Non-blocking memory transfer: {self.use_non_blocking}")
log.info("----------------------")
log.info("-" * 25)
def forward_vace(
self,
@@ -2318,6 +2319,8 @@ class WanModel(torch.nn.Module):
sdancer_input=None, # SteadyDancer
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
scail_input=None, # SCAIL pose
dual_control_input=None, # LongVie2 dual controlnet
transformer_options={},
):
r"""
Forward pass through the diffusion model
@@ -2544,6 +2547,16 @@ class WanModel(torch.nn.Module):
x = [u.flatten(2).transpose(1, 2) for u in x]
self.original_seq_len = x[0].shape[1]
prev_latent = None
if dual_control_input is not None:
prev_latent = dual_control_input.get("prev_latent", None)
if prev_latent is not None:
F += prev_latent.shape[2]
prev_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in prev_latent]
prev_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in prev_x]
seq_len += prev_x[0].shape[1]
x = [torch.cat([u, v], dim=1) for u, v in zip(prev_x, x)]
# SCAIL pose
if scail_input is not None:
scail_pose_latents = scail_input.get("pose_latent", None)
@@ -2834,6 +2847,44 @@ class WanModel(torch.nn.Module):
chunked_self_attention = False
seq_chunks = 0
# dual control
if dual_control_input is not None and dual_control_input["start_percent"] <= current_step_percentage <= dual_control_input["end_percent"]:
dense_latent = dual_control_input["dense_input_latent"]
print("dense_latent shape:", dense_latent.shape)
sparse_latent = dual_control_input["sparse_input_latent"]
if dense_latent is None and sparse_latent is None:
raise ValueError("At least one of dense_input_latent or sparse_input_latent must be provided in dual_control_input")
if dense_latent is not None:
dense_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in dense_latent]
dense_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in dense_x]
dense = self.dual_controller.control_initial_combine_linear_dense(dense_x[0])
if sparse_latent is not None:
sparse_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in sparse_latent]
sparse_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in sparse_x]
sparse = self.dual_controller.control_initial_combine_linear_sparse(sparse_x[0])
if dense_latent is None:
dense = torch.zeros_like(sparse)
elif sparse_latent is None:
sparse = torch.zeros_like(dense)
control_context = clip_fea_control = None
if context != []:
control_context = self.dual_controller.control_text_linear(context)
if clip_embed is not None:
clip_fea_control = self.dual_controller.control_text_linear(clip_embed)
control_t_mod = self.dual_controller.control_t_mod(e0)
control_freqs = torch.cat([
self.dual_controller_freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
self.dual_controller_freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
self.dual_controller_freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
else:
dual_control_input = None
# MultiTalk
if multitalk_audio is not None:
self.multitalk_audio_proj.to(self.main_device)
@@ -3039,6 +3090,7 @@ class WanModel(torch.nn.Module):
camera_embed=camera_embed,
audio_proj=audio_proj,
num_latent_frames = F,
frame_tokens=x.shape[1] // F,
original_seq_len=self.original_seq_len,
enhance_enabled=enhance_enabled,
audio_scale=audio_scale,
@@ -3066,6 +3118,7 @@ class WanModel(torch.nn.Module):
e_tr=e0_token_replace if use_token_replace else None,
tr_start=token_replace_start,
tr_num=replace_token_num,
transformer_options=transformer_options
)
if self.audio_model is not None:
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
@@ -3125,8 +3178,22 @@ class WanModel(torch.nn.Module):
if lynx_ref_buffer is None and lynx_ref_feature_extractor:
lynx_ref_buffer = {}
attn_override_blocks = attention_mode = None
attention_mode_override_active = False
attention_mode_override = transformer_options.get("attention_mode_override", None)
if attention_mode_override is not None:
attn_override_blocks = attention_mode_override.get("blocks", range(len(self.blocks)))
if attention_mode_override["start_step"] <= current_step < attention_mode_override["end_step"]:
attention_mode_override_active = True
if attention_mode_override["verbose"]:
tqdm.write(f"Applying attention mode override: {attention_mode_override['mode']} at step {current_step} on blocks: {attn_override_blocks if attn_override_blocks is not None else 'all'}")
for b, block in enumerate(self.blocks):
mm.throw_exception_if_processing_interrupted()
if attention_mode_override_active and b in attn_override_blocks:
attention_mode = attention_mode_override['mode']
else:
attention_mode = None
block_idx = f"{b:02d}"
if lynx_ref_buffer is not None and not lynx_ref_feature_extractor:
lynx_ref_feature = lynx_ref_buffer.get(block_idx, None)
@@ -3170,9 +3237,21 @@ class WanModel(torch.nn.Module):
x_onetoall_ref = onetoall_ref_block_samples[b // interval_ref]
# ---run block----#
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, **kwargs)
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, attention_mode_override=attention_mode, **kwargs)
# ---post block----#
# dual controlnet
if dual_control_input is not None and (hasattr(block, "control_blocks_dense") or hasattr(block, "control_blocks_sparse")):
if dense_latent is not None and hasattr(block, "control_blocks_dense"):
dense = block.control_blocks_dense(dense, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
if sparse_latent is not None and hasattr(block, "control_blocks_sparse"):
sparse = block.control_blocks_sparse(sparse, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
if prev_latent is not None:
x[:, -self.original_seq_len:] += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
else:
x += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
if self.audio_injector is not None and s2v_audio_input is not None:
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
if block.has_face_fuser_block and motion_vec is not None:
@@ -3275,8 +3354,10 @@ class WanModel(torch.nn.Module):
# x = x[:, :self.original_seq_len]
#grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
x = x[:, :self.original_seq_len]
if prev_latent is not None:
x = x[:, -self.original_seq_len:]
else:
x = x[:, :self.original_seq_len]
x = self.head(x, e.to(x.device), temp_length=F,
e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)
+1 -1
View File
@@ -34,7 +34,7 @@ class ERSDEScheduler():
sigmas.append(0.0)
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas * self.num_train_timesteps
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
self.step_index = 0
self.old_denoised = None
self.old_denoised_d = None
@@ -35,9 +35,7 @@ class FlowMatchSchedulerResMultistep():
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas * self.num_train_timesteps
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
def step(self, model_output, timestep, sample):
if timestep.ndim == 2:
@@ -48,14 +46,14 @@ class FlowMatchSchedulerResMultistep():
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.tensor(0)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
x0_pred = (sample - sigma * model_output)
if sigma_next == 0 or self.prev_model_output is None:
@@ -73,7 +71,7 @@ class FlowMatchSchedulerResMultistep():
self.old_sigma_next = sigma_next
self.prev_model_output = x0_pred
return x
def add_noise(self, original_samples, noise, timestep):
"""
+42 -33
View File
@@ -989,7 +989,8 @@ class VideoVAE_(nn.Module):
mean=None,
inv_std=None,
pruning_rate=0.0,
cpu_cache=False):
cpu_cache=False,
verbose=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -1000,6 +1001,7 @@ class VideoVAE_(nn.Module):
self.temperal_upsample = temperal_downsample[::-1]
self.mean = mean
self.inv_std = inv_std
self.verbose = verbose
# modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
@@ -1085,12 +1087,13 @@ class VideoVAE_(nn.Module):
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
eps = torch.randn_like(std)
return mu + std * eps
try:
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
if self.verbose:
try:
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return mu
@@ -1137,7 +1140,7 @@ class VideoVAE_(nn.Module):
except:
pass
x = self.conv2(z)
for i in range(iter_):
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
self._conv_idx = [0]
if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :],
@@ -1154,12 +1157,13 @@ class VideoVAE_(nn.Module):
if pbar:
pbar.update_absolute(0)
self.clear_cache()
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return out
def reparameterize(self, mu, log_var):
@@ -1186,12 +1190,12 @@ class VideoVAE_(nn.Module):
class WanVideoVAE(nn.Module):
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False):
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
super().__init__()
self.dtype = dtype
self.cpu_cache = cpu_cache
self.verbose = verbose
mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
@@ -1205,7 +1209,7 @@ class WanVideoVAE(nn.Module):
self.z_dim = z_dim
# init model
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache).eval().requires_grad_(False)
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
self.upsampling_factor = 8
@@ -1431,7 +1435,8 @@ class VideoVAE38_(VideoVAE_):
mean=None,
inv_std=None,
pruning_rate=0.0,
cpu_cache=False):
cpu_cache=False,
verbose=False):
super(VideoVAE_, self).__init__()
self.dim = dim
self.z_dim = z_dim
@@ -1444,6 +1449,7 @@ class VideoVAE38_(VideoVAE_):
self.mean = mean
self.inv_std = inv_std
self.cpu_cache = cpu_cache
self.verbose = verbose
# modules
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
@@ -1481,12 +1487,13 @@ class VideoVAE38_(VideoVAE_):
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
self.clear_cache()
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return mu
@@ -1519,18 +1526,19 @@ class VideoVAE38_(VideoVAE_):
pbar.update(1)
out = unpatchify(out, patch_size=2)
self.clear_cache()
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return out
class WanVideoVAE38(WanVideoVAE):
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False):
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
super(WanVideoVAE, self).__init__()
mean = [
@@ -1554,7 +1562,8 @@ class WanVideoVAE38(WanVideoVAE):
self.dtype = dtype
self.z_dim = z_dim
self.cpu_cache = cpu_cache
self.verbose = verbose
# init model
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache).eval().requires_grad_(False)
self.upsampling_factor = 16
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
self.upsampling_factor = 16