diff --git a/nodes.py b/nodes.py index dd30935..7a6bc4d 100644 --- a/nodes.py +++ b/nodes.py @@ -396,7 +396,10 @@ class WanVideoModelLoader: in_channels = sd["patch_embedding.weight"].shape[1] print("in_channels: ", in_channels) ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0] - model_type = "i2v" if in_channels == 36 else "t2v" + if in_channels in [36, 48]: + model_type = "i2v" + elif in_channels == 16: + model_type = "t2v" num_heads = 40 if dim == 5120 else 12 num_layers = 40 if dim == 5120 else 30 @@ -409,7 +412,13 @@ class WanVideoModelLoader: "i2v_720": [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], } if model_type == "i2v": - model_variant = "i2v_480" if "480" in model else "i2v_720" + if "480" in model: + model_variant = "i2v_480" + elif "720" in model: + model_variant = "i2v_720" + else: + model_variant = "1_3B" + elif model_type == "t2v": model_variant = "14B" if dim == 5120 else "1_3B" log.info(f"Model variant detected: {model_variant}") @@ -467,6 +476,8 @@ class WanVideoModelLoader: patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher.model.is_patched = False + + control_lora = False if lora is not None: for l in lora: @@ -481,8 +492,10 @@ class WanVideoModelLoader: #spacepxl's control LoRA patch # for key in lora_sd.keys(): # print(key) + if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd: log.info("Control-LoRA detected, patching model...") + control_lora = True in_cls = transformer.patch_embedding.__class__ # nn.Conv3d old_in_dim = transformer.in_dim # 16 @@ -645,6 +658,7 @@ class WanVideoModelLoader: patcher.model["manual_offloading"] = manual_offloading patcher.model["quantization"] = "disabled" patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False + patcher.model["control_lora"] = control_lora if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} @@ -1336,6 +1350,7 @@ class WanVideoImageToVideoEncode: }, "optional": { "end_image": ("IMAGE", {"tooltip": "end frame"}), + "control_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "Control signal for the Fun -model"}), } } @@ -1345,7 +1360,7 @@ class WanVideoImageToVideoEncode: CATEGORY = "WanVideoWrapper" def process(self, vae, start_image, width, height, num_frames, clip_embeds, force_offload, noise_aug_strength, - start_latent_strength, end_latent_strength, end_image=None): + start_latent_strength, end_latent_strength, end_image=None, control_embeds=None): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1397,7 +1412,10 @@ class WanVideoImageToVideoEncode: concatenated = torch.cat([resized_image.to(device), zero_frames, resized_end_image.to(device)], dim=1) * end_latent_strength y = vae.encode([concatenated.to(device=device, dtype=vae.dtype)], device, end_=(end_image is not None))[0] - y = torch.cat([mask, y]) + if control_embeds is None: + y = torch.cat([mask, y]) + else: + y[:, 1:] = 0 # Calculate maximum sequence length patches_per_frame = lat_h * lat_w // (patch_size[1] * patch_size[2]) @@ -1418,6 +1436,7 @@ class WanVideoImageToVideoEncode: "num_frames": num_frames, "lat_h": lat_h, "lat_w": lat_w, + "control_embeds": control_embeds, "end_image": resized_end_image if end_image is not None else None } @@ -1675,6 +1694,8 @@ class WanVideoSampler: model = model.model transformer = model.diffusion_model + control_lora = model["control_lora"] + device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1743,7 +1764,14 @@ class WanVideoSampler: image_cond = image_embeds.get("image_embeds", None) print("image_cond", image_cond.shape) clip_fea = image_embeds.get("clip_context", None) - clip_fea_neg = image_embeds.get("negative_clip_context", None) + clip_fea_neg = image_embeds.get("negative_clip_context", None) + + control_embeds = image_embeds.get("control_embeds", None) + if control_embeds is not None: + control_latents = control_embeds.get("control_images", None) + control_start_percent = control_embeds.get("start_percent", 0.0) + control_end_percent = control_embeds.get("end_percent", 1.0) + else: #t2v target_shape = image_embeds.get("target_shape", None) if target_shape is None: @@ -2009,21 +2037,30 @@ class WanVideoSampler: nonlocal patcher current_step_percentage = idx / len(timesteps) - control_enabled = False + control_lora_enabled = False if control_latents is not None: - control_enabled = True - if not control_start_percent <= current_step_percentage <= control_end_percent: - image_cond = None - control_enabled = False - if patcher.model.is_patched: - log.info("Unloading LoRA...") - patcher.unpatch_model(device) - patcher.model.is_patched = False + if control_lora: + control_lora_enabled = True else: - if not patcher.model.is_patched: - log.info("Loading LoRA...") - patcher = apply_lora(patcher, device, device, low_mem_load=False) - patcher.model.is_patched = True + if control_start_percent <= current_step_percentage <= control_end_percent: + image_cond = torch.cat([control_latents, image_cond]) + else: + image_cond = torch.cat([torch.zeros_like(image_cond), image_cond]) + + if not control_start_percent <= current_step_percentage <= control_end_percent: + if control_lora: + image_cond = None + control_lora_enabled = False + if patcher.model.is_patched: + log.info("Unloading LoRA...") + patcher.unpatch_model(device) + patcher.model.is_patched = False + else: + if control_lora: + if not patcher.model.is_patched: + log.info("Loading LoRA...") + patcher = apply_lora(patcher, device, device, low_mem_load=False) + patcher.model.is_patched = True base_params = { 'seq_len': seq_len, @@ -2032,7 +2069,7 @@ class WanVideoSampler: 't': timestep, 'current_step': idx, 'y': [image_cond] if image_cond is not None else None, - 'control_enabled': control_enabled, + 'control_lora_enabled': control_lora_enabled, } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index bf0c642..4ec7783 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -812,7 +812,7 @@ class WanModel(ModelMixin, ConfigMixin): freqs=None, current_step=0, pred_id=None, - control_enabled=False, + control_lora_enabled=False, ): r""" Forward pass through the diffusion model @@ -848,7 +848,7 @@ class WanModel(ModelMixin, ConfigMixin): x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] # embeddings - if control_enabled: + if control_lora_enabled: self.expanded_patch_embedding.to(device) x = [ self.expanded_patch_embedding(u.unsqueeze(0))