Refactor Fun control input
Allows using reference only
This commit is contained in:
+2221
-2031
File diff suppressed because it is too large
Load Diff
@@ -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"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user