From 36249e1a973b89a4e528d62520e190d83b6fe1c5 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 5 Apr 2025 15:53:56 +0300 Subject: [PATCH] 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. --- nodes.py | 25 +++++++++++++++++++++---- 1 file changed, 21 insertions(+), 4 deletions(-) diff --git a/nodes.py b/nodes.py index caa958d..6a65921 100644 --- a/nodes.py +++ b/nodes.py @@ -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)