Basic Skyreels A2 support

converted model here:

https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Wan2_1_SkyreelsA2_fp8_e4m3fn.safetensors

Encode the reference and input to the ImageToVideoEncode -node, take into account the extra latent when choosing frame count (should be -4 to normal), otherwise the standard I2V workflow.
This commit is contained in:
kijai
2025-04-05 15:53:56 +03:00
parent ae40c3c131
commit 36249e1a97
+21 -4
View File
@@ -1426,6 +1426,7 @@ class WanVideoImageToVideoEncode:
"control_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "Control signal for the Fun -model"}),
"fun_model": ("BOOLEAN", {"default": False, "tooltip": "Enable when using Fun model"}),
"temporal_mask": ("MASK", {"tooltip": "mask"}),
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
}
}
@@ -1435,7 +1436,7 @@ class WanVideoImageToVideoEncode:
CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, clip_embeds, force_offload, noise_aug_strength,
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_model=False, temporal_mask=None):
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_model=False, temporal_mask=None, extra_latents=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -1514,6 +1515,13 @@ class WanVideoImageToVideoEncode:
concatenated = resized_start_image[:,:num_frames] * temporal_mask[:num_frames].unsqueeze(0)
y = vae.encode([concatenated.to(device=device, dtype=vae.dtype)], device, end_=(end_image is not None and not fun_model))[0]
drop_last = False
if extra_latents is not None:
samples = extra_latents["samples"].squeeze(0)
y = torch.cat([y, samples], dim=1)
mask = torch.cat([mask, torch.ones_like(mask[:, 0:1])], dim=1)
num_frames += samples.shape[1] * 4
drop_last = True
y[:, :1] *= start_latent_strength
y[:, -1:] *= end_latent_strength
if control_embeds is None:
@@ -1547,7 +1555,8 @@ class WanVideoImageToVideoEncode:
"lat_w": lat_w,
"control_embeds": control_embeds["control_embeds"] if control_embeds is not None else None,
"end_image": resized_end_image if end_image is not None else None,
"fun_model": fun_model
"fun_model": fun_model,
"drop_last": drop_last,
}
return (image_embeds,)
@@ -2071,7 +2080,7 @@ class WanVideoSampler:
control_latents, clip_fea, clip_fea_neg, end_image = None, None, None, None
vace_data, vace_context, vace_scale = None, None, None
fun_model, has_ref = False, False
fun_model, has_ref, drop_last = False, False, False
image_cond = image_embeds.get("image_embeds", None)
@@ -2103,6 +2112,7 @@ class WanVideoSampler:
control_latents = control_embeds["control_images"].to(device)
control_start_percent = control_embeds.get("start_percent", 0.0)
control_end_percent = control_embeds.get("end_percent", 1.0)
drop_last = image_embeds.get("drop_last", False)
else: #t2v
target_shape = image_embeds.get("target_shape", None)
if target_shape is None:
@@ -2821,7 +2831,7 @@ class WanVideoSampler:
pass
return ({
"samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_model else None, "has_ref": has_ref
"samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_model else None, "has_ref": has_ref, "drop_last": drop_last,
}, )
class WindowTracker:
@@ -2888,8 +2898,11 @@ class WanVideoDecode:
latents = samples["samples"]
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
drop_last = samples.get("drop_last", False)
is_looped = samples.get("looped", False)
print("drop_last", drop_last)
vae.to(device)
latents = latents.to(device = device, dtype = vae.dtype)
@@ -2899,6 +2912,10 @@ class WanVideoDecode:
if has_ref:
latents = latents[:, :, 1:]
if drop_last:
print("latents.shape", latents.shape)
latents = latents[:, :, :-1]
print("latents.shape", latents.shape)
#if is_looped:
# latents = torch.cat([latents[:, :, :warmup_latent_count],latents], dim=2)