Lynx + VACE

This commit is contained in:
kijai
2025-09-27 02:05:17 +03:00
parent cf235c0728
commit 358baa2c2a
3 changed files with 46 additions and 28 deletions
+11 -7
View File
@@ -30,8 +30,8 @@ class WanLynxIPCrossAttention(nn.Module):
b, n, d = x.size(0), block.num_heads, block.head_dim
if self.registers is not None:
print("self.registers.shape", self.registers.shape) #torch.Size([1, 16, 5120])
print("ip_x.shape", ip_x.shape) #torch.Size([1, 16, 5120])
#print("self.registers.shape", self.registers.shape) #torch.Size([1, 16, 5120])
#print("ip_x.shape", ip_x.shape) #torch.Size([1, 16, 5120])
ip_lens = [ip_x.shape[1]]
ip_x_list = vector_to_list(ip_x, ip_lens, 1)
@@ -39,14 +39,18 @@ class WanLynxIPCrossAttention(nn.Module):
ip_x, ip_lens = list_to_vector(ip_x_list, 1)
ip_key = self.to_k_ip(ip_x)
ip_key = ip_key * torch.rsqrt(ip_key.pow(2).mean(dim=-1, keepdim=True) + 1e-5).to(ip_key.dtype)
ip_value = self.to_v_ip(ip_x)
ip_key = ip_key.view(b, -1, n, d)
ip_value = ip_value.view(b, -1, n, d)
if self.registers is None: # lite model normalization
ip_key = ip_key * torch.rsqrt(ip_key.pow(2).mean(dim=-1, keepdim=True) + 1e-5).to(ip_key.dtype)
else: # full model normalization
ip_key = block.norm_k(ip_key)
ip_x = attention(q, ip_key, ip_value).reshape(b, -1, n * d)
ip_x = attention(
q,
ip_key.view(b, -1, n, d),
ip_value.view(b, -1, n, d)
).reshape(b, -1, n * d)
return ip_x
+27 -18
View File
@@ -561,19 +561,24 @@ class WanVideoExtraModelSelect:
"required": {
"extra_model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' path to extra state dict to add to the main model"}),
},
"optional": {
"prev_model":("VACEPATH", {"default": None, "tooltip": "For loading multiple extra models"}),
},
}
RETURN_TYPES = ("VACEPATH",)
RETURN_NAMES = ("extra_model", )
FUNCTION = "getvacepath"
FUNCTION = "getmodelpath"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Extra model to load and add to the main model, ie. VACE or MTV Crafter 'ComfyUI/models/diffusion_models'"
def getvacepath(self, extra_model):
extra_model = {
"path": folder_paths.get_full_path("diffusion_models", extra_model),
}
return (extra_model,)
def getmodelpath(self, extra_model, prev_model=None):
extra_model = {"path": folder_paths.get_full_path("diffusion_models", extra_model)}
if prev_model is not None and isinstance(prev_model, list):
extra_model_list = prev_model + [extra_model]
else:
extra_model_list = [extra_model]
return (extra_model_list,)
class WanVideoLoraBlockEdit:
def __init__(self):
@@ -1102,18 +1107,20 @@ class WanVideoModelLoader:
# currently this can be VAE or MTV-Crafter weights
if extra_model is not None:
if gguf:
if not extra_model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the extra model must also be GGUF quantized, if the main model already has VACE included, you can disconnect the extra module loader")
extra_sd, extra_reader = load_gguf(extra_model["path"])
gguf_reader.append(extra_reader)
del extra_reader
else:
if extra_model["path"].endswith(".gguf"):
raise ValueError("With GGUF extra model the main model must also be GGUF quantized model")
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
sd.update(extra_sd)
del extra_sd
for _model in extra_model:
print("Loading extra model: ", _model["path"])
if gguf:
if not _model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the extra model must also be GGUF quantized, if the main model already has VACE included, you can disconnect the extra module loader")
extra_sd, extra_reader = load_gguf(_model["path"])
gguf_reader.append(extra_reader)
del extra_reader
else:
if _model["path"].endswith(".gguf"):
raise ValueError("With GGUF extra model the main model must also be GGUF quantized model")
extra_sd = load_torch_file(_model["path"], device=transformer_load_device, safe_load=True)
sd.update(extra_sd)
del extra_sd
first_key = next(iter(sd))
if first_key.startswith("model.diffusion_model."):
@@ -1144,9 +1151,11 @@ class WanVideoModelLoader:
#lynx
lynx_layers = "none"
if "blocks.0.cross_attn.ip_adapter.to_v_ip.weight" in sd and "blocks.0.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"
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"
+8 -3
View File
@@ -1307,11 +1307,15 @@ class BaseWanAttentionBlock(WanAttentionBlock):
cross_attn_norm=False,
eps=1e-6,
block_id=None,
block_idx=0,
attention_mode='sdpa',
rope_func="comfy",
rms_norm_function="default"
rms_norm_function="default",
lynx_layers="none"
):
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode, rope_func, rms_norm_function=rms_norm_function)
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm,
cross_attn_norm, eps, attention_mode, rope_func, rms_norm_function=rms_norm_function,
block_idx=block_idx, lynx_layers=lynx_layers)
self.block_id = block_id
def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs):
@@ -1674,7 +1678,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)
block_id=self.vace_layers_mapping[i] if i in self.vace_layers else None, lynx_layers=lynx_layers, block_idx=i)
for i in range(num_layers)
])
else:
@@ -2113,6 +2117,7 @@ class WanModel(torch.nn.Module):
submodule.step = current_step
lynx_x_ip = None
lynx_ip_scale = 1.0
if lynx_embeds is not None:
if not is_uncond:
lynx_x_ip = lynx_embeds["ip_x"].to(self.main_device)