remove loras if SetLora bypassed/disconnected

This commit is contained in:
kijai
2025-07-23 17:51:59 +03:00
parent 66226a8a1b
commit bdce322e65
3 changed files with 17 additions and 6 deletions
+8 -2
View File
@@ -107,6 +107,12 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N
has_scale = hasattr(submodule, "scale_weight")
has_lora = hasattr(submodule, "lora")
if has_scale or has_lora:
original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward)
original_forward_ = submodule.forward
setattr(submodule, "original_forward_", original_forward_)
setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input))
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
if hasattr(submodule, "lora"):
delattr(submodule, "lora")
+5 -2
View File
@@ -8,7 +8,7 @@ import hashlib
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from .wanvideo.modules.model import rope_params
from .fp8_optimization import convert_linear_with_lora_and_scale
from .fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module
from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list
from .multitalk.multitalk import timestep_transform, add_noise
@@ -1281,9 +1281,12 @@ class WanVideoSampler:
control_lora = model["control_lora"]
transformer_options = patcher.model_options.get("transformer_options", None)
remove_lora_from_module(transformer)
if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True:
log.info(f"Using {len(patcher.patches)} patches for WanVideo model")
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
else:
remove_lora_from_module(transformer)
#compile
compile_args = model["compile_args"]
+4 -2
View File
@@ -629,8 +629,10 @@ class WanVideoSetLoRAs:
"required":
{
"model": ("WANVIDEOMODEL", ),
"lora": ("WANVIDLORA", ),
},
"optional": {
"lora": ("WANVIDLORA", ),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
@@ -640,7 +642,7 @@ class WanVideoSetLoRAs:
EXPERIMENTAL = True
DESCRIPTION = "Sets the LoRA weights to be used directly in linear layers of the model, this does NOT merge LoRAs"
def setlora(self, model, lora):
def setlora(self, model, lora=None):
if lora is None:
return (model,)