diff --git a/lynx/modules.py b/lynx/modules.py index ec9e8ae..1dd0b8f 100644 --- a/lynx/modules.py +++ b/lynx/modules.py @@ -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 diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 324ccd2..affb7ce 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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" diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3f446fa..f57cdb1 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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)