Lynx + VACE
This commit is contained in:
+11
-7
@@ -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
@@ -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"
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user