Allow selecting which blocks to use for lynx ref
This commit is contained in:
+26
-2
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user