diff --git a/nodes.py b/nodes.py index 31f3528..e1b0d03 100644 --- a/nodes.py +++ b/nodes.py @@ -25,6 +25,7 @@ from comfy.utils import load_torch_file, save_torch_file, ProgressBar, common_up import comfy.model_base import comfy.latent_formats from comfy.clip_vision import clip_preprocess, ClipVisionModel +from comfy.sd import load_lora_for_models script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -432,9 +433,10 @@ class WanVideoModelLoader: comfy_model.load_device = transformer_load_device patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) + patcher.is_patched = False if lora is not None: - from comfy.sd import load_lora_for_models + for l in lora: log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") lora_path = l["path"] @@ -444,10 +446,41 @@ class WanVideoModelLoader: if l["blocks"]: lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"]) + #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...") + + in_cls = transformer.patch_embedding.__class__ # nn.Conv3d + old_in_dim = transformer.in_dim # 16 + new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1] + assert new_in_dim == 32 + + new_in = in_cls( + new_in_dim, + transformer.patch_embedding.out_channels, + transformer.patch_embedding.kernel_size, + transformer.patch_embedding.stride, + transformer.patch_embedding.padding, + ).to(device=device, dtype=torch.bfloat16) + + new_in.weight.zero_() + new_in.bias.zero_() + + new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight) + new_in.bias.copy_(transformer.patch_embedding.bias) + + transformer.patch_embedding = new_in + transformer.expanded_patch_embedding = new_in + transformer.register_to_config(in_dim=new_in_dim) + patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) + del lora_sd patcher.patch_model(device) + patcher.is_patched = True del sd gc.collect() @@ -992,6 +1025,40 @@ class WanVideoEmptyEmbeds: } return (embeds,) + +class WanVideoControlEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the control signal"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the control signal"}), + }, + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, latents, start_percent, end_percent): + + samples = latents["samples"].squeeze(0) + C, T, H, W = samples.shape + + num_frames = (T - 1) * 4 + 1 + seq_len = math.ceil((H * W) / 4 * ((num_frames - 1) // 4 + 1)) + + embeds = { + "max_seq_len": seq_len, + "target_shape": samples.shape, + "num_frames": num_frames, + "control_images": samples, + "start_percent": start_percent, + "end_percent": end_percent, + } + + return (embeds,) #region Sampler @@ -1176,6 +1243,17 @@ class WanVideoSampler: device=torch.device("cpu"), generator=seed_g) + control_latents = image_embeds.get("control_images", None) + if control_latents is not None: + image_cond = control_latents.to(device) + control_start_percent = image_embeds.get("start_percent", 0.0) + control_end_percent = image_embeds.get("end_percent", 1.0) + + if not patcher.is_patched: + print("Patching model for control") + patcher.patch_model(device) + patcher.is_patched = True + latent_video_length = noise.shape[1] if context_options is not None: @@ -1373,6 +1451,17 @@ class WanVideoSampler: def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, teacache_state=None): with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): + + current_step_percentage = idx / len(timesteps) + if not control_start_percent <= current_step_percentage <= control_end_percent: + image_cond = None + control_enabled = False + if patcher.is_patched: + patcher.unpatch_model(device) + patcher.is_patched = False + else: + control_enabled = True + base_params = { 'clip_fea': clip_fea, 'seq_len': seq_len, @@ -1381,6 +1470,7 @@ class WanVideoSampler: 't': timestep, 'current_step': idx, 'y': image_cond, + 'control_enabled': control_enabled, } # Get conditional prediction @@ -1589,7 +1679,10 @@ class WanVideoSampler: partial_img_emb = None if image_cond is not None: partial_img_emb = image_cond[:, c, :, :] - partial_img_emb[:, 0, :, :] = image_cond[:, 0, :, :].to(intermediate_device) + partial_image_cond = image_cond[:, 0, :, :].to(intermediate_device) + if min(c) > len(c) // 2: + partial_image_cond *= 0.1 + partial_img_emb[:, 0, :, :] = partial_image_cond partial_latent_model_input = latent_model_input[:, c, :, :] @@ -1801,7 +1894,7 @@ class WanVideoEncode: mask = torch.nn.functional.interpolate( mask.unsqueeze(1), # Add channel dim for interpolate size=(target_h, target_w), - mode='nearest' + mode='bilinear' ).squeeze(1) # Remove channel dim # Add batch & channel dims for final output @@ -1913,6 +2006,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoVRAMManagement": WanVideoVRAMManagement, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoFlowEdit": WanVideoFlowEdit, + "WanVideoControlEmbeds": WanVideoControlEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -1937,4 +2031,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoVRAMManagement": "WanVideo VRAM Management", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoFlowEdit": "WanVideo FlowEdit", + "WanVideoControlEmbeds": "WanVideo Control Embeds", } diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 43caf3d..cf057d4 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -520,6 +520,10 @@ class WanModel(ModelMixin, ConfigMixin): # embeddings self.patch_embedding = nn.Conv3d( in_dim, dim, kernel_size=patch_size, stride=patch_size) + + self.original_patch_embedding = self.patch_embedding + self.expanded_patch_embedding = self.patch_embedding + self.text_embedding = nn.Sequential( nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'), nn.Linear(dim, dim)) @@ -591,7 +595,8 @@ class WanModel(ModelMixin, ConfigMixin): device=torch.device('cuda'), freqs=None, current_step=0, - pred_id=None + pred_id=None, + control_enabled=False, ): r""" Forward pass through the diffusion model @@ -625,7 +630,11 @@ class WanModel(ModelMixin, ConfigMixin): x = torch.cat([x, y], dim=0) # embeddings - x = [self.patch_embedding(x.unsqueeze(0))] + if control_enabled: + x = [self.expanded_patch_embedding(x.unsqueeze(0))] + else: + x = [self.original_patch_embedding(x.unsqueeze(0))] + grid_sizes = torch.stack( [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) x = [u.flatten(2).transpose(1, 2) for u in x]