From 260c181e0ee11e0fcd1595c430c3c5357870e2c6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 7 Oct 2025 15:07:31 +0300 Subject: [PATCH] Allow lynx reference adapter to run with I2V models --- nodes_sampler.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index 37e5143..1b23125 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1023,8 +1023,14 @@ class WanVideoSampler: lynx_embeds["ref_feature_extractor"] = True log.info(f"Lynx ref latent shape: {lynx_ref_latent.shape}") log.info("Extracting Lynx ref cond buffer...") + if transformer.in_dim == 36: + mask_latents = torch.tile(torch.zeros_like(lynx_ref_latent[:1]), [4, 1, 1, 1]) + empty_image_cond = torch.cat([mask_latents, torch.zeros_like(lynx_ref_latent)], dim=0).to(device) + lynx_ref_input = torch.cat([lynx_ref_latent, empty_image_cond], dim=0) + else: + lynx_ref_input = lynx_ref_latent lynx_ref_buffer = transformer( - [lynx_ref_latent.to(device, dtype)], + [lynx_ref_input.to(device, dtype)], torch.tensor([0], device=device), lynx_ref_text_embed["prompt_embeds"], seq_len=math.ceil((lynx_ref_latent.shape[2] * lynx_ref_latent.shape[3]) / 4 * lynx_ref_latent.shape[1]), @@ -1033,8 +1039,12 @@ class WanVideoSampler: log.info(f"Extracted {len(lynx_ref_buffer)} cond ref buffers") if not math.isclose(cfg[0], 1.0): log.info("Extracting Lynx ref uncond buffer...") + if transformer.in_dim == 36: + lynx_ref_input_uncond = torch.cat([lynx_ref_latent_uncond, empty_image_cond], dim=0) + else: + lynx_ref_input_uncond = lynx_ref_latent_uncond lynx_ref_buffer_uncond = transformer( - [lynx_ref_latent_uncond.to(device, dtype)], + [lynx_ref_input_uncond.to(device, dtype)], torch.tensor([0], device=device), lynx_ref_text_embed["prompt_embeds"], seq_len=math.ceil((lynx_ref_latent.shape[2] * lynx_ref_latent.shape[3]) / 4 * lynx_ref_latent.shape[1]),