From 073092955800033911031a983e693b07f43dae5a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 1 Oct 2025 19:04:56 +0300 Subject: [PATCH] Allow selecting which blocks to use for lynx ref --- lynx/nodes.py | 28 ++++++++++++++++++++++++++-- wanvideo/modules/model.py | 10 +++++++--- 2 files changed, 33 insertions(+), 5 deletions(-) diff --git a/lynx/nodes.py b/lynx/nodes.py index 1e25755..289f120 100644 --- a/lynx/nodes.py +++ b/lynx/nodes.py @@ -206,7 +206,8 @@ class WanVideoAddLynxEmbeds: "vae": ("WANVAE", {"tooltip": "VAE model, only needed if ref_image is provided"}), "lynx_ip_embeds": ("LYNXIP", {"tooltip": "lynx face embeddings"}), "ref_image": ("IMAGE",), - "ref_text_embed": ("WANVIDEOTEXTEMBEDS",) + "ref_text_embed": ("WANVIDEOTEXTEMBEDS",), + "ref_blocks_to_use": ("STRING", {"default": "", "forceInput": True, "tooltip": "Comma-separated list of block indices and ranges to use for reference feature, e.g. '0-20, 25, 28, 35-39'. If empty, use all blocks."}), } } @@ -215,7 +216,7 @@ class WanVideoAddLynxEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, ip_scale, ref_scale, start_percent, end_percent, lynx_cfg_scale, vae=None, lynx_ip_embeds=None, ref_image=None, ref_text_embed=None): + def add(self, embeds, ip_scale, ref_scale, start_percent, end_percent, lynx_cfg_scale, vae=None, lynx_ip_embeds=None, ref_image=None, ref_text_embed=None, ref_blocks_to_use=""): if ref_image is not None and ref_text_embed is None: raise ValueError("If ref_image is provided, ref_text_embed must also be provided.") if ref_image is not None: @@ -224,6 +225,28 @@ class WanVideoAddLynxEmbeds: ref_latent = vae.encode([ref_image_in], device, tiled=False, sample=True) ref_latent_uncond = vae.encode([torch.zeros_like(ref_image_in)], device, tiled=False, sample=True) vae.to(offload_device) + if ref_blocks_to_use.strip() == "": + ref_blocks_to_use = None + else: + # Parse comma-separated blocks and ranges + blocks = [] + for item in ref_blocks_to_use.split(","): + item = item.strip() + if "-" in item and not item.startswith("-"): + # Handle range like "0-20" or "35-39" + try: + start, end = item.split("-", 1) + start, end = int(start.strip()), int(end.strip()) + blocks.extend(list(range(start, end + 1))) + except ValueError: + print(f"Invalid range format: {item}") + elif item.isdigit(): + # Handle single number + blocks.append(int(item)) + else: + print(f"Invalid block specification: {item}") + ref_blocks_to_use = sorted(list(set(blocks))) # Remove duplicates and sort + print("Using ref blocks:", ref_blocks_to_use) new_entry = { "ip_x": lynx_ip_embeds["ip_x"] if lynx_ip_embeds is not None else None, @@ -236,6 +259,7 @@ class WanVideoAddLynxEmbeds: "cfg_scale": lynx_cfg_scale, "start_percent": start_percent, "end_percent": end_percent, + "ref_blocks_to_use": ref_blocks_to_use, } updated = dict(embeds) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3c388b6..fcea61b 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2102,6 +2102,8 @@ class WanModel(torch.nn.Module): lynx_ip_scale = lynx_ref_scale = 1.0 if lynx_embeds is not None: lynx_ref_feature_extractor = lynx_embeds.get("ref_feature_extractor", False) + lynx_ref_blocks_to_use = lynx_embeds.get("ref_blocks_to_use", [list(range(len(self.blocks)))]) + print(f"Using Lynx ref feature extractor: {lynx_ref_feature_extractor}, blocks: {lynx_ref_blocks_to_use}") if (lynx_embeds['start_percent'] <= current_step_percentage <= lynx_embeds['end_percent']) and not lynx_ref_feature_extractor: if not is_uncond: lynx_x_ip = lynx_embeds.get("ip_x", None) @@ -2659,8 +2661,9 @@ class WanModel(torch.nn.Module): for b, block in enumerate(self.blocks): block_idx = f"{b:02d}" if lynx_ref_buffer is not None and not lynx_ref_feature_extractor: - #print("reading from lynx ref buffer for block", block_idx) lynx_ref_feature = lynx_ref_buffer.get(block_idx, None) + if lynx_ref_feature is not None: + print("loading from lynx ref buffer for block", block_idx) else: lynx_ref_feature = None # Prefetch blocks if enabled @@ -2707,8 +2710,9 @@ class WanModel(torch.nn.Module): log.info(f"Block {b}: transfer_time={transfer_time:.4f}s, compute_time={compute_time:.4f}s, to_cpu_transfer_time={to_cpu_transfer_time:.4f}s") # lynx ref if lynx_ref_feature_extractor: - print("storing to lynx ref buffer for block", block_idx) - lynx_ref_buffer[block_idx] = lynx_ref_feature + if b in lynx_ref_blocks_to_use: + print("storing to lynx ref buffer for block", block_idx) + lynx_ref_buffer[block_idx] = lynx_ref_feature #uni3c controlnet if uni3c_controlnet_states is not None and b < len(uni3c_controlnet_states): x[:, :self.original_seq_len] += uni3c_controlnet_states[b].to(x) * uni3c_data["controlnet_weight"]