Add start/end percent for Phantom and other small fixes

This commit is contained in:
kijai
2025-04-23 01:00:41 +03:00
parent 2b0bf44994
commit e0e5fcf713
2 changed files with 33 additions and 10 deletions
+30 -7
View File
@@ -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,
+3 -3
View File
@@ -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)