Allow selecting which blocks to use for lynx ref

This commit is contained in:
kijai
2025-10-01 19:04:56 +03:00
parent 1e88104558
commit 0730929558
2 changed files with 33 additions and 5 deletions
+26 -2
View File
@@ -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)
+7 -3
View File
@@ -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"]