From bb23263f7825c84e8173743a553c8e0faecc97a9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 6 Oct 2025 18:42:14 +0300 Subject: [PATCH] Separate full model's ip and ref layer loading Allows using full ref with lite ip adapter, or full ref alone without loading the ip weights --- lynx/modules.py | 9 ++++----- lynx/nodes.py | 32 ++------------------------------ nodes_model_loading.py | 20 +++++++++++--------- nodes_sampler.py | 21 +++++++++++++-------- wanvideo/modules/model.py | 37 ++++++++++++++++++------------------- 5 files changed, 48 insertions(+), 71 deletions(-) diff --git a/lynx/modules.py b/lynx/modules.py index 89b922f..2a80f9e 100644 --- a/lynx/modules.py +++ b/lynx/modules.py @@ -29,9 +29,8 @@ class WanLynxIPCrossAttention(nn.Module): else: self.registers = None - def forward(self, block, q, x, ip_x): - b, n, d = x.size(0), block.num_heads, block.head_dim - s = q.shape[1] + def forward(self, block, q, ip_x): + b, s, n, d = q.shape ip_lens = [ip_x.shape[1]] if self.registers is not None and ip_x is not None and ip_x.shape[0] == 1: @@ -54,10 +53,10 @@ class WanLynxIPCrossAttention(nn.Module): q, ip_key.view(b, -1, n, d), ip_value.view(b, -1, n, d) - ).reshape(b, -1, n * d) + ).flatten(2) -@torch.compiler.disable() +#@torch.compiler.disable() class WanLynxRefAttention(nn.Module): def __init__(self, dim=5120, bias=True, attention_mode="sdpa"): super().__init__() diff --git a/lynx/nodes.py b/lynx/nodes.py index 289f120..919b8bf 100644 --- a/lynx/nodes.py +++ b/lynx/nodes.py @@ -1,10 +1,7 @@ import os import torch -import gc -from ..utils import log, dict_to_device +from ..utils import log import numpy as np -from accelerate import init_empty_weights -from accelerate.utils import set_module_tensor_to_device import comfy.model_management as mm from comfy.utils import load_torch_file @@ -56,19 +53,6 @@ class LoadLynxResampler: return resampler, -class VideoStyleInfo: # key names should match those used in style.yaml file - style_name: str = 'none' - num_frames: int = 81 - seed: int = -1 - guidance_scale: float = 5.0 - guidance_scale_i: float = 2.0 - num_inference_steps: int = 50 - width: int = 832 - height: int = 480 - prompt: str = '' - negative_prompt: str = '' - - class LynxInsightFaceCrop: @classmethod def INPUT_TYPES(s): @@ -89,19 +73,8 @@ class LynxInsightFaceCrop: from insightface.utils import face_align image_np = (image[0].numpy() * 255).astype(np.uint8) - # Landmarks - - # landmarks = np.array([ - # [599.9878, 633.5308 ], - # [893.5392, 642.8297 ], - # [733.05945, 844.21454], - # [640.42206, 970.78687], - # [854.4849, 979.5486 ] - # ]) landmarks = get_landmarks_from_image(image_np) - print(landmarks) - in_image = np.array(image_np) landmark = np.array(landmarks) @@ -137,12 +110,11 @@ class LynxEncodeFaceIP: from .face.face_encoder import FaceEncoderArcFace image_in = ip_image.permute(0, 3, 1, 2).to(device) * 2 - 1 # to [-1, 1] - print("image_in.shape", image_in.shape) #torch.Size([1, 3, 112, 112]) # Face embedding via ArcFace face_encoder = FaceEncoderArcFace() face_encoder.init_encoder_model(device) - arcface_embed = face_encoder(image_in).to(device, resampler.dtype) + arcface_embed = face_encoder(image_in).to(device, resampler.dtype)[0] arcface_embed = arcface_embed.reshape([1, -1, 512]) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 0ec2d42..192795f 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1152,15 +1152,16 @@ class WanVideoModelLoader: is_wananimate = "pose_patch_embedding.weight" in sd #lynx - lynx_layers = "none" - if "blocks.0.cross_attn.ip_adapter.to_v_ip.weight" in sd and "blocks.0.self_attn.ref_adapter.to_k_ref.weight" in sd: - log.info("Lynx full model detected") - n_registers = sd["blocks.0.cross_attn.ip_adapter.registers"].shape[1] - lynx_layers = "full" + lynx_ip_layers = lynx_ref_layers = None + if "blocks.0.self_attn.ref_adapter.to_k_ref.weight" in sd: + log.info("Lynx full reference adapter detected") + lynx_ref_layers = "full" + if "blocks.0.cross_attn.ip_adapter.registers" in sd: + log.info("Lynx full IP adapter detected") + lynx_ip_layers = "full" elif "blocks.0.cross_attn.ip_adapter.to_v_ip.weight" in sd: - log.info("Lynx lite model detected") - n_registers = 0 - lynx_layers = "lite" + log.info("Lynx lite IP adapter detected") + lynx_ip_layers = "lite" model_type = "t2v" if "audio_injector.injector.0.k.weight" in sd: @@ -1298,7 +1299,8 @@ class WanVideoModelLoader: "humo_audio": is_humo, "is_wananimate": is_wananimate, "rms_norm_function": rms_norm_function, - "lynx_layers": lynx_layers, + "lynx_ip_layers": lynx_ip_layers, + "lynx_ref_layers": lynx_ref_layers, } diff --git a/nodes_sampler.py b/nodes_sampler.py index 83e0be5..9440a0f 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1003,6 +1003,9 @@ class WanVideoSampler: lynx_ref_buffer = None lynx_embeds = image_embeds.get("lynx_embeds", None) if lynx_embeds is not None: + if lynx_embeds.get("ip_x", None) is not None: + if transformer.blocks[0].cross_attn.ip_adapter is None: + raise ValueError("Lynx IP embeds provided, but the no lynx ip adapter layers found in the model.") lynx_embeds = lynx_embeds.copy() log.info("Using Lynx embeddings", lynx_embeds) lynx_ref_latent = lynx_embeds.get("ref_latent", None) @@ -1013,6 +1016,8 @@ class WanVideoSampler: lynx_cfg_scale = [lynx_cfg_scale] * (steps + 1) if lynx_ref_latent is not None: + if transformer.blocks[0].self_attn.ref_adapter is None: + raise ValueError("Lynx reference provided, but the no lynx reference adapter layers found in the model.") lynx_ref_latent = lynx_ref_latent[0] lynx_ref_latent_uncond = lynx_ref_latent_uncond[0] lynx_embeds["ref_feature_extractor"] = True @@ -1038,14 +1043,14 @@ class WanVideoSampler: ) log.info(f"Extracted {len(lynx_ref_buffer_uncond)} uncond ref buffers") - if lynx_embeds.get("ip_x", None) is not None: - lynx_embeds["ip_x"] = lynx_embeds["ip_x"].to(device, dtype) - lynx_embeds["ip_x_uncond"] = lynx_embeds["ip_x_uncond"].to(device, dtype) - lynx_embeds["ref_feature_extractor"] = False - lynx_embeds["ref_latent"] = lynx_embeds["ref_text_embed"] = None - lynx_embeds["ref_buffer"] = lynx_ref_buffer - lynx_embeds["ref_buffer_uncond"] = lynx_ref_buffer_uncond if not math.isclose(cfg[0], 1.0) else None - mm.soft_empty_cache() + if lynx_embeds.get("ip_x", None) is not None: + lynx_embeds["ip_x"] = lynx_embeds["ip_x"].to(device, dtype) + lynx_embeds["ip_x_uncond"] = lynx_embeds["ip_x_uncond"].to(device, dtype) + lynx_embeds["ref_feature_extractor"] = False + lynx_embeds["ref_latent"] = lynx_embeds["ref_text_embed"] = None + lynx_embeds["ref_buffer"] = lynx_ref_buffer + lynx_embeds["ref_buffer_uncond"] = lynx_ref_buffer_uncond if not math.isclose(cfg[0], 1.0) else None + mm.soft_empty_cache() #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index fdd4562..26e2801 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -640,7 +640,7 @@ class WanT2VCrossAttention(WanSelfAttention): q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d) if nag_context is not None and not is_uncond: - x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) + x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: k = self.norm_k(self.k(context)).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d) @@ -650,13 +650,10 @@ class WanT2VCrossAttention(WanSelfAttention): q = rope_apply_z(q, grid_sizes, cross_freqs, inner_t).to(q) k = rope_apply_c(k, cross_freqs, inner_c).to(q) - x_text = attention(q, k, v, attention_mode=self.attention_mode) - x_text = x_text.flatten(2) - - x = x_text + x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2) if lynx_x_ip is not None and self.ip_adapter is not None and ip_scale !=0: - lynx_x_ip = self.ip_adapter(self, q, x, lynx_x_ip) + lynx_x_ip = self.ip_adapter(self, q, lynx_x_ip) x = x.add(lynx_x_ip, alpha=lynx_ip_scale) # FantasyTalking audio attention @@ -855,7 +852,8 @@ class WanAttentionBlock(nn.Module): use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, - lynx_layers="none", + lynx_ip_layers=None, + lynx_ref_layers=None, block_idx=0 ): super().__init__() @@ -908,11 +906,13 @@ class WanAttentionBlock(nn.Module): # Lynx self.ref_adapter = None - if lynx_layers == "full": - from ...lynx.modules import WanLynxIPCrossAttention, WanLynxRefAttention - self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=self.dim, dim=self.dim, n_registers=16) + if lynx_ref_layers == "full": + from ...lynx.modules import WanLynxRefAttention self.self_attn.ref_adapter = WanLynxRefAttention(dim=self.dim) - elif lynx_layers == "lite": + if lynx_ip_layers == "full": + from ...lynx.modules import WanLynxIPCrossAttention + self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=self.dim, dim=self.dim, n_registers=16) + elif lynx_ip_layers == "lite": from ...lynx.modules import WanLynxIPCrossAttention if self.block_idx % 2 == 0: self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=2048, dim=self.dim, n_registers=0, bias=False) @@ -1525,7 +1525,8 @@ class WanModel(torch.nn.Module): is_wananimate=False, motion_encoder_dim=512, # lynx - lynx_layers="none", + lynx_ip_layers=None, + lynx_ref_layers=None, ): r""" Initialize the diffusion model backbone. @@ -1635,7 +1636,8 @@ class WanModel(torch.nn.Module): self.multitalk_model_type = "none" - self.lynx_layers = lynx_layers + self.lynx_ip_layers = lynx_ip_layers + self.lynx_ref_layers = lynx_ref_layers self.humo_audio = humo_audio @@ -1680,7 +1682,7 @@ class WanModel(torch.nn.Module): BaseWanAttentionBlock('t2v_cross_attn', self.in_features, self.out_features, ffn_dim, self.ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function, - block_id=self.vace_layers_mapping[i] if i in self.vace_layers else None, lynx_layers=lynx_layers, block_idx=i) + block_id=self.vace_layers_mapping[i] if i in self.vace_layers else None, lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers, block_idx=i) for i in range(num_layers) ]) else: @@ -1697,7 +1699,7 @@ class WanModel(torch.nn.Module): qk_norm, cross_attn_norm, eps, attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function, use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio, - face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_layers=lynx_layers, block_idx=i) + face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers, block_idx=i) for i in range(num_layers) ]) #MTV Crafter @@ -2105,7 +2107,6 @@ class WanModel(torch.nn.Module): lynx_ref_blocks_to_use = lynx_embeds.get("ref_blocks_to_use", None) if lynx_ref_blocks_to_use is None: lynx_ref_blocks_to_use = list(range(len(self.blocks))) - print(f"Using Lynx ref feature extractor: {lynx_ref_feature_extractor}, blocks: {lynx_ref_blocks_to_use}") if (lynx_embeds['start_percent'] <= current_step_percentage <= lynx_embeds['end_percent']) and not lynx_ref_feature_extractor: if not is_uncond: lynx_x_ip = lynx_embeds.get("ip_x", None) @@ -2671,8 +2672,6 @@ class WanModel(torch.nn.Module): 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) - if lynx_ref_feature is not None: - print("loading from lynx ref buffer for block", block_idx) else: lynx_ref_feature = None # Prefetch blocks if enabled @@ -2720,7 +2719,7 @@ class WanModel(torch.nn.Module): # lynx ref if lynx_ref_feature_extractor: if b in lynx_ref_blocks_to_use: - print("storing to lynx ref buffer for block", block_idx) + log.info(f"storing to lynx ref buffer for block {block_idx}") lynx_ref_buffer[block_idx] = lynx_ref_feature #uni3c controlnet if uni3c_controlnet_states is not None and b < len(uni3c_controlnet_states):