From e0e5fcf7133c49f9ba589e6ad5632ae99e5c751d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 23 Apr 2025 01:00:41 +0300 Subject: [PATCH] Add start/end percent for Phantom and other small fixes --- nodes.py | 37 ++++++++++++++++++++++++++++++------- wanvideo/modules/model.py | 6 +++--- 2 files changed, 33 insertions(+), 10 deletions(-) diff --git a/nodes.py b/nodes.py index 52b40a7..102fe06 100644 --- a/nodes.py +++ b/nodes.py @@ -1759,11 +1759,14 @@ class WanVideoPhantomEmbeds: "phantom_latent_1": ("LATENT", {"tooltip": "reference latents for the phantom model"}), "phantom_cfg_scale": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the extra phantom cond pass"}), + "phantom_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the phantom model"}), + "phantom_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the phantom model"}), }, "optional": { "phantom_latent_2": ("LATENT", {"tooltip": "reference latents for the phantom model"}), "phantom_latent_3": ("LATENT", {"tooltip": "reference latents for the phantom model"}), "phantom_latent_4": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + #"vace_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "VACE embeds"}), } } @@ -1772,7 +1775,7 @@ class WanVideoPhantomEmbeds: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, num_frames, phantom_cfg_scale, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None): + def process(self, num_frames, phantom_cfg_scale, phantom_start_percent, phantom_end_percent, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None, vace_embeds=None): vae_stride = (4, 8, 8) samples = phantom_latent_1["samples"].squeeze(0) if phantom_latent_2 is not None: @@ -1792,7 +1795,20 @@ class WanVideoPhantomEmbeds: "num_frames": num_frames, "phantom_latents": samples, "phantom_cfg_scale": phantom_cfg_scale, + "phantom_start_percent": phantom_start_percent, + "phantom_end_percent": phantom_end_percent, } + if vace_embeds is not None: + vace_input = { + "vace_context": vace_embeds["vace_context"], + "vace_scale": vace_embeds["vace_scale"], + "has_ref": vace_embeds["has_ref"], + "vace_start_percent": vace_embeds["vace_start_percent"], + "vace_end_percent": vace_embeds["vace_end_percent"], + "vace_seq_len": vace_embeds["vace_seq_len"], + "additional_vace_inputs": vace_embeds["additional_vace_inputs"], + } + embeds.update(vace_input) return (embeds,) @@ -2431,6 +2447,8 @@ class WanVideoSampler: phantom_latents = image_embeds.get("phantom_latents", None) phantom_cfg_scale = image_embeds.get("phantom_cfg_scale", None) + phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0) + phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) if phantom_latents is not None: phantom_latents = phantom_latents.to(device) @@ -2741,14 +2759,19 @@ class WanVideoSampler: image_cond_input = image_cond z = z.to(dtype) - z_pos = z_neg = z if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) - if phantom_latents is not None: - z_pos = torch.cat([z_pos[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) - z_phantom_img = torch.cat([z_pos[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) - z_neg = torch.cat([z_pos[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1) + use_phantom = False + if phantom_latents is not None and \ + (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \ + (phantom_end_percent > 0 and idx == 0 and current_step_percentage >= phantom_start_percent): + z_pos = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_phantom_img = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_neg = torch.cat([z[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1) + use_phantom = True + else: + z_pos = z_neg = z base_params = { 'seq_len': seq_len, @@ -2795,7 +2818,7 @@ class WanVideoSampler: ) noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device) #phantom - if phantom_latents is not None: + if use_phantom: noise_pred_phantom, teacache_state_phantom = transformer( [z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, y=[image_cond_input] if image_cond_input is not None else None, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 034aafd..cceccbb 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -958,7 +958,7 @@ class WanModel(ModelMixin, ConfigMixin): kwargs ): # embeddings - c = [self.vace_patch_embedding(u.unsqueeze(0)) for u in vace_context] + c = [self.vace_patch_embedding(u.unsqueeze(0).float()).to(x.dtype) for u in vace_context] c = [u.flatten(2).transpose(1, 2) for u in c] c = torch.cat([ torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], @@ -1218,11 +1218,11 @@ class WanModel(ModelMixin, ConfigMixin): if (data["start"] <= current_step_percentage <= data["end"]) or \ (data["end"] > 0 and current_step == 0 and current_step_percentage >= data["start"]): - vace_hints = self.forward_vace(x.to(torch.float32), data["context"], data["seq_len"], kwargs) + vace_hints = self.forward_vace(x, data["context"], data["seq_len"], kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(data["scale"]) else: - vace_hints = self.forward_vace(x.to(torch.float32), vace_data, seq_len, kwargs) + vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(1.0)