diff --git a/LongVie2/modules.py b/LongVie2/modules.py new file mode 100644 index 0000000..7d89789 --- /dev/null +++ b/LongVie2/modules.py @@ -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) diff --git a/LongVie2/nodes.py b/LongVie2/nodes.py new file mode 100644 index 0000000..d686874 --- /dev/null +++ b/LongVie2/nodes.py @@ -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", + } diff --git a/__init__.py b/__init__.py index 4efb858..3dbdcb3 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] \ No newline at end of file +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/custom_linear.py b/custom_linear.py index b80be20..9fe4762 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -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 diff --git a/multitalk/multitalk_loop.py b/multitalk/multitalk_loop.py index ed8b0af..513edcc 100644 --- a/multitalk/multitalk_loop.py +++ b/multitalk/multitalk_loop.py @@ -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() diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 1f9a1d8..8e28176 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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" } diff --git a/nodes_sampler.py b/nodes_sampler.py index 5eabb25..cdeab61 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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) diff --git a/ultravico/sageattn/attn_qk_int8_per_block.py b/ultravico/sageattn/attn_qk_int8_per_block.py index 3d85856..645d5a4 100644 --- a/ultravico/sageattn/attn_qk_int8_per_block.py +++ b/ultravico/sageattn/attn_qk_int8_per_block.py @@ -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) diff --git a/utils.py b/utils.py index 243e5a6..4636dee 100644 --- a/utils.py +++ b/utils.py @@ -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/ diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 09caac8..f31f6be 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -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 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 96df41c..9a84f6f 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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) diff --git a/wanvideo/schedulers/ersde_scheduler.py b/wanvideo/schedulers/ersde_scheduler.py index 6abd7a7..9574dfb 100644 --- a/wanvideo/schedulers/ersde_scheduler.py +++ b/wanvideo/schedulers/ersde_scheduler.py @@ -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 diff --git a/wanvideo/schedulers/flowmatch_res_multistep.py b/wanvideo/schedulers/flowmatch_res_multistep.py index 68eb3bd..5c6fb99 100644 --- a/wanvideo/schedulers/flowmatch_res_multistep.py +++ b/wanvideo/schedulers/flowmatch_res_multistep.py @@ -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): """ diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 757a77f..e957bac 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -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 \ No newline at end of file + 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