From 1a1af029129762ac1b2a9d5b656fdcc8c9abae59 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 23 Jul 2025 12:51:57 +0300 Subject: [PATCH] Allow using different init image for context windows beyond first one This can be useful with models like MAGREF that behave differently when the init image is padded with white, idea is to use full init for first window and padded image for the rest, so that it works more like reference and doesn't force the window to snap back to the init. --- nodes.py | 32 +++++++++++++++++++++++++------- 1 file changed, 25 insertions(+), 7 deletions(-) diff --git a/nodes.py b/nodes.py index b2ad9ef..0e8145b 100644 --- a/nodes.py +++ b/nodes.py @@ -1036,6 +1036,7 @@ class WanVideoContextOptions: }, "optional": { "fuse_method": (["linear", "pyramid"], {"default": "linear", "tooltip": "Window weight function: linear=ramps at edges only, pyramid=triangular weights peaking in middle"}), + "reference_latent": ("LATENT", {"tooltip": "Image to be used as init for I2V models for windows where first frame is not the actual first frame. Mostly useful with MAGREF model"}), } } @@ -1045,7 +1046,7 @@ class WanVideoContextOptions: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Context options for WanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow." - def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None, fuse_method="linear"): + def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise, verbose, image_cond_start_step=6, image_cond_window_count=2, vae=None, fuse_method="linear", reference_latent=None): context_options = { "context_schedule":context_schedule, "context_frames":context_frames, @@ -1053,7 +1054,8 @@ class WanVideoContextOptions: "context_overlap":context_overlap, "freenoise":freenoise, "verbose":verbose, - "fuse_method":fuse_method + "fuse_method":fuse_method, + "reference_latent":reference_latent["samples"][0] if reference_latent is not None else None, } return (context_options,) @@ -1536,14 +1538,16 @@ class WanVideoSampler: # Context windows is_looped = False - if context_options is not None: + context_reference_latent = None + if context_options is not None: context_schedule = context_options["context_schedule"] context_frames = (context_options["context_frames"] - 1) // 4 + 1 context_stride = context_options["context_stride"] // 4 context_overlap = context_options["context_overlap"] // 4 - context_vae = context_options.get("vae", None) - if context_vae is not None: - context_vae.to(device) + # context_vae = context_options.get("vae", None) + # if context_vae is not None: + # context_vae.to(device) + context_reference_latent = context_options.get("reference_latent", None) # Get total number of prompts num_prompts = len(text_embeds["prompt_embeds"]) @@ -2408,7 +2412,21 @@ class WanVideoSampler: partial_control_latents = None if image_cond is not None: partial_img_emb = image_cond[:, c] - partial_img_emb[:, 0] = image_cond[:, 0].to(intermediate_device) + + if c[0] != 0 and context_reference_latent is not None: + new_init_image = context_reference_latent[:, 0].to(intermediate_device) + # Concatenate the first 4 channels of partial_img_emb with new_init_image to match the required shape + if new_init_image.shape[0] + 4 == partial_img_emb.shape[0]: + partial_img_emb[:, 0] = torch.cat([ + image_cond[:4, 0], + new_init_image + ], dim=0) + else: + # fallback to original assignment if shape matches + partial_img_emb[:, 0] = new_init_image + else: + new_init_image = image_cond[:, 0].to(intermediate_device) + partial_img_emb[:, 0] = new_init_image if control_latents is not None: partial_control_latents = control_latents[:, c]