Allow further filtering of LoRA layers to apply

This commit is contained in:
kijai
2025-05-30 16:24:44 +03:00
parent ef40577b70
commit 5c676d3a49
+33 -15
View File
@@ -182,23 +182,36 @@ class WanVideoModelConfig:
self.latent_format.latent_channels = 16
self.manual_cast_dtype = dtype
self.sampling_settings = {"multiplier": 1.0}
# Don't know what this is. Value taken from ComfyUI Mochi model.
self.memory_usage_factor = 2.0
# denoiser is handled by extension
self.unet_config["disable_unet_model_creation"] = True
def filter_state_dict_by_blocks(state_dict, blocks_mapping):
def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]):
filtered_dict = {}
print("layer_filter: ", layer_filter)
print("blocks_mapping: ", blocks_mapping)
if isinstance(layer_filter, str):
layer_filters = [layer_filter] if layer_filter else []
else:
# Filter out empty strings
layer_filters = [f for f in layer_filter if f] if layer_filter else []
print("layer_filter: ", layer_filters)
for key in state_dict:
if 'blocks.' in key:
block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2]
block_key = f'{block_pattern[0]}.{block_pattern[1]}.'
if not any(filter_str in key for filter_str in layer_filters):
if 'blocks.' in key:
block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2]
block_key = f'{block_pattern[0]}.{block_pattern[1]}.'
if block_key in blocks_mapping:
if block_key in blocks_mapping:
filtered_dict[key] = state_dict[key]
else:
filtered_dict[key] = state_dict[key]
else:
filtered_dict[key] = state_dict[key]
for key in filtered_dict:
print(key)
return filtered_dict
@@ -392,14 +405,15 @@ class WanVideoLoraSelect:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras"
def getlorapath(self, lora, strength, blocks=None, prev_lora=None, low_mem_load=False):
def getlorapath(self, lora, strength, blocks={}, prev_lora=None, low_mem_load=False):
loras_list = []
lora = {
"path": folder_paths.get_full_path("loras", lora),
"strength": strength,
"name": lora.split(".")[0],
"blocks": blocks,
"blocks": blocks.get("selected_blocks", {}),
"layer_filter": blocks.get("layer_filter", ""),
"low_mem_load": low_mem_load,
}
if prev_lora is not None:
@@ -441,7 +455,7 @@ class WanVideoLoraBlockEdit:
for i in range(40):
arg_dict["blocks.{}.".format(i)] = argument
return {"required": arg_dict}
return {"required": arg_dict, "optional": {"layer_filter": ("STRING", {"default": "", "multiline": True})}}
RETURN_TYPES = ("SELECTEDBLOCKS", )
RETURN_NAMES = ("blocks", )
@@ -450,10 +464,14 @@ class WanVideoLoraBlockEdit:
CATEGORY = "WanVideoWrapper"
def select(self, **kwargs):
def select(self, layer_filter=[], **kwargs):
selected_blocks = {k: v for k, v in kwargs.items() if v is True and isinstance(v, bool)}
print("Selected blocks LoRA: ", selected_blocks)
return (selected_blocks,)
selected = {
"selected_blocks": selected_blocks,
"layer_filter": [x.strip() for x in layer_filter.split(",")]
}
return (selected,)
#region Model loading
class WanVideoModelLoader:
@@ -718,7 +736,7 @@ class WanVideoModelLoader:
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
#spacepxl's control LoRA patch
# for key in lora_sd.keys():