Refactor Fun control input

Allows using reference only
This commit is contained in:
kijai
2025-08-14 13:03:51 +03:00
parent 64a0c03673
commit 5691ae5314
2 changed files with 2292 additions and 2056 deletions
+71 -25
View File
@@ -905,13 +905,13 @@ class WanVideoImageToVideoEncode:
has_ref = True
y[:, :1] *= start_latent_strength
y[:, -1:] *= end_latent_strength
if control_embeds is None:
y = torch.cat([mask, y])
else:
if end_image is None:
y[:, 1:] = 0
elif start_image is None:
y[:, -1:] = 0
# if control_embeds is None:
# y = torch.cat([mask, y])
# else:
# if end_image is None:
# y[:, 1:] = 0
# elif start_image is None:
# y[:, -1:] = 0
# Calculate maximum sequence length
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
@@ -940,7 +940,7 @@ class WanVideoImageToVideoEncode:
"fun_or_fl2v_model": fun_or_fl2v_model,
"has_ref": has_ref,
"add_cond_latents": add_cond_latents,
"mask": mask if control_embeds is not None else None, # for 2.2 Fun control as it can handle masks
"mask": mask
}
return (image_embeds,)
@@ -1115,11 +1115,11 @@ 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"}),
},
"optional": {
"latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}),
"fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}),
}
}
@@ -1129,10 +1129,10 @@ class WanVideoControlEmbeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, latents, start_percent, end_percent, fun_ref_image=None):
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
def process(self, start_percent, end_percent, fun_ref_image=None, latents=None):
if latents is not None:
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))
@@ -1151,6 +1151,38 @@ class WanVideoControlEmbeds:
return (embeds,)
class WanVideoAddControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"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"}),
},
"optional": {
"latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}),
"fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, embeds, start_percent, end_percent, fun_ref_image=None, latents=None):
new_entry = {
"control_images": latents["samples"].squeeze(0) if latents is not None else None,
"start_percent": start_percent,
"end_percent": end_percent,
"fun_ref_image": fun_ref_image["samples"][:,:, 0] if fun_ref_image is not None else None,
}
updated = dict(embeds)
updated["control_embeds"] = new_entry
return (updated,)
class WanVideoSLG:
@classmethod
def INPUT_TYPES(s):
@@ -1658,7 +1690,16 @@ class WanVideoSampler:
if image_cond is not None:
if transformer.in_dim == 16:
raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models")
if transformer.in_dim not in [48, 32]: # fun 2.1 models don't use the mask
image_cond_mask = image_embeds.get("mask", None)
if image_cond_mask is not None:
image_cond = torch.cat([image_cond_mask, image_cond])
else:
image_cond[:, 1:] = 0
log.info(f"image_cond shape: {image_cond.shape}")
#ATI tracks
if transformer_options is not None:
ATI_tracks = transformer_options.get("ati_tracks", None)
@@ -1701,12 +1742,13 @@ class WanVideoSampler:
control_embeds = image_embeds.get("control_embeds", None)
if control_embeds is not None:
print("control_embeds:", control_embeds)
if transformer.in_dim not in [52, 48, 36, 32]:
raise ValueError("Control signal only works with Fun-Control model")
if transformer.in_dim == 52 or transformer.control_adapter is not None: #fun 2.2 control
image_cond_mask = image_embeds.get("mask", None)
if image_cond_mask is not None:
image_cond = torch.cat([image_cond_mask, image_cond])
# if transformer.in_dim == 52 or transformer.control_adapter is not None: #fun 2.2 control
# image_cond_mask = image_embeds.get("mask", None)
# if image_cond_mask is not None:
# image_cond = torch.cat([image_cond_mask, image_cond])
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)
@@ -1723,6 +1765,8 @@ class WanVideoSampler:
raise ValueError("Empty image embeds must be provided for T2V models")
has_ref = image_embeds.get("has_ref", False)
# VACE
vace_context = image_embeds.get("vace_context", None)
vace_scale = image_embeds.get("vace_scale", None)
if not isinstance(vace_scale, list):
@@ -1776,6 +1820,7 @@ class WanVideoSampler:
log.info(f"RecamMaster source video shape: {recam_latents.shape}")
seq_len *= 2
# Fun control and control lora
control_embeds = image_embeds.get("control_embeds", None)
if control_embeds is not None:
control_latents = control_embeds.get("control_images", None)
@@ -1785,8 +1830,6 @@ class WanVideoSampler:
control_camera_latents = control_embeds.get("control_camera_latents", None)
control_camera_start_percent = control_embeds.get("control_camera_start_percent", 0.0)
control_camera_end_percent = control_embeds.get("control_camera_end_percent", 1.0)
if control_camera_latents is not None:
control_camera_latents = control_camera_latents.to(device)
if control_lora:
image_cond = control_latents.to(device)
@@ -1819,6 +1862,7 @@ class WanVideoSampler:
masked_video_latents_input = torch.zeros_like(noise)
image_cond = torch.cat([mask_latents, masked_video_latents_input], dim=0).to(device)
# Phantom inputs
phantom_latents = image_embeds.get("phantom_latents", None)
phantom_cfg_scale = image_embeds.get("phantom_cfg_scale", None)
if not isinstance(phantom_cfg_scale, list):
@@ -2238,20 +2282,20 @@ class WanVideoSampler:
current_step_percentage = idx / len(timesteps)
control_lora_enabled = False
image_cond_input = None
if control_latents is not None:
if control_embeds is not None and control_camera_latents is None:
if control_lora:
control_lora_enabled = True
else:
if (control_start_percent <= current_step_percentage <= control_end_percent) or \
(control_end_percent > 0 and idx == 0 and current_step_percentage >= control_start_percent):
if ((control_start_percent <= current_step_percentage <= control_end_percent) or \
(control_end_percent > 0 and idx == 0 and current_step_percentage >= control_start_percent)) and \
(control_latents is not None):
image_cond_input = torch.cat([control_latents.to(z), image_cond.to(z)])
else:
image_cond_input = torch.cat([torch.zeros_like(control_latents, dtype=dtype), image_cond.to(z)])
image_cond_input = torch.cat([torch.zeros_like(noise, device=device, dtype=dtype), image_cond.to(z)])
if fun_ref_image is not None:
fun_ref_input = fun_ref_image.to(z)
else:
fun_ref_input = torch.zeros_like(z, dtype=z.dtype)[:, 0].unsqueeze(1)
#fun_ref_input = None
if control_lora:
if not control_start_percent <= current_step_percentage <= control_end_percent:
@@ -3500,6 +3544,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoLatentReScale": WanVideoLatentReScale,
"WanVideoScheduler": WanVideoScheduler,
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -3531,5 +3576,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoTextEncodeCached": "WanVideo TextEncode Cached",
"WanVideoAddExtraLatent": "WanVideo Add Extra Latent",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent"
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent",
"WanVideoAddControlEmbeds": "WanVideo Add Control Embeds"
}