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
This commit is contained in:
kijai
2025-10-06 18:42:14 +03:00
parent 174cba5759
commit bb23263f78
5 changed files with 48 additions and 71 deletions
+4 -5
View File
@@ -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__()
+2 -30
View File
@@ -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])
+11 -9
View File
@@ -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,
}
+13 -8
View File
@@ -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,
+18 -19
View File
@@ -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):