From 45e39e758de2f18ff7b1f78442f82690dbc5f41c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 4 Apr 2025 02:26:24 +0300 Subject: [PATCH] Allow selecting VACE blocks --- nodes.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/nodes.py b/nodes.py index b60a743..851e281 100644 --- a/nodes.py +++ b/nodes.py @@ -294,18 +294,22 @@ class WanVideoVACEModelSelect: return { "required": { "vace_model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' VACE model to use when not using model that has it included"}), + "vace_blocks": ("STRING", {"default": "0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28", "multiline": True, "tooltip": "Blocks to apply VACE to, default is for 1.3B model"}), }, } RETURN_TYPES = ("VACEPATH",) - RETURN_NAMES = ("vace_path", ) + RETURN_NAMES = ("vace_model", ) FUNCTION = "getvacepath" CATEGORY = "WanVideoWrapper" DESCRIPTION = "VACE model to use when not using model that has it included, loaded from 'ComfyUI/models/diffusion_models'" - def getvacepath(self, vace_model): - vace_model_path = folder_paths.get_full_path_or_raise("diffusion_models", vace_model) - return (vace_model_path,) + def getvacepath(self, vace_model, vace_blocks): + vace_model = { + "path": folder_paths.get_full_path("diffusion_models", vace_model), + "blocks": [int(x.strip()) for x in vace_blocks.split(",")], + } + return (vace_model,) class WanVideoLoraBlockEdit: def __init__(self): @@ -412,7 +416,7 @@ class WanVideoModelLoader: sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True) if vace_model is not None: - vace_sd = load_torch_file(vace_model, device=transformer_load_device, safe_load=True) + vace_sd = load_torch_file(vace_model["path"], device=transformer_load_device, safe_load=True) sd.update(vace_sd) first_key = next(iter(sd)) @@ -437,7 +441,10 @@ class WanVideoModelLoader: vace_layers, vace_in_dim = None, None if "vace_blocks.0.after_proj.weight" in sd: model_type = "t2v" - vace_layers = [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28] + if vace_model is not None: + vace_layers = vace_model["blocks"] + else: + vace_layers = [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 28] vace_in_dim = 96 log.info(f"Model type: {model_type}, num_heads: {num_heads}, num_layers: {num_layers}")