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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user