remove loras if SetLora bypassed/disconnected
This commit is contained in:
+8
-2
@@ -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")
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user