Allow LoRA timestep scheduling

LoRA strength can now be a list of floats that represents the strength at any given step
This commit is contained in:
kijai
2025-08-04 00:02:19 +03:00
parent 88b5feb1ed
commit 801f543a7a
4 changed files with 28 additions and 9 deletions
+7 -2
View File
@@ -30,8 +30,12 @@ def fp8_linear_forward(cls, original_dtype, input):
@torch.compiler.disable()
def apply_lora(weight, lora):
def apply_lora(weight, lora, step=None):
for lora_diff, lora_strength in zip(lora[0], lora[1]):
if isinstance(lora_strength, list):
lora_strength = lora_strength[step]
if lora_strength == 0.0:
return weight
patch_diff = torch.mm(
lora_diff[0].flatten(start_dim=1).to(weight.device),
lora_diff[1].flatten(start_dim=1).to(weight.device)
@@ -57,7 +61,7 @@ def linear_with_lora_and_scale_forward(cls, input):
lora = getattr(cls, "lora", None)
if lora is not None:
weight = apply_lora(weight, lora).to(input.dtype)
weight = apply_lora(weight, lora, cls.step).to(input.dtype)
return torch.nn.functional.linear(input, weight, bias)
@@ -110,6 +114,7 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N
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))
setattr(submodule, "step", 0) # Initialize step for LoRA if needed
def remove_lora_from_module(module):
+2
View File
@@ -1468,6 +1468,8 @@ class WanVideoSampler:
log.info("Unloading all LoRAs")
remove_lora_from_module(transformer)
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
#torch.compile
if model["auto_cpu_offload"] is False:
transformer = compile_model(transformer, model["compile_args"])
+12 -6
View File
@@ -371,11 +371,12 @@ class WanVideoLoraSelect:
low_mem_load = False # Unmerged LoRAs don't need low_mem_load
loras_list = []
strength = round(strength, 4)
if strength == 0.0:
if prev_lora is not None:
loras_list.extend(prev_lora)
return (loras_list,)
if not isinstance(strength, list):
strength = round(strength, 4)
if strength == 0.0:
if prev_lora is not None:
loras_list.extend(prev_lora)
return (loras_list,)
try:
lora_path = folder_paths.get_full_path("loras", lora)
@@ -666,11 +667,14 @@ class WanVideoSetLoRAs:
merge_loras = l.get("merge_loras", True)
if merge_loras is True:
raise ValueError("Set LoRA node does not use low_mem_load and can't merge LoRAs, disable 'merge_loras' in the LoRA select node.")
patcher.model_options['transformer_options']["lora_scheduling_enabled"] = False
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
if isinstance(lora_strength, list):
patcher.model_options['transformer_options']["lora_scheduling_enabled"] = True
if lora_strength == 0:
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
continue
@@ -1073,6 +1077,8 @@ class WanVideoModelLoader:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
if isinstance(lora_strength, list):
transformer.lora_scheduling_enabled = True
if lora_strength == 0:
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
continue
+7 -1
View File
@@ -1096,6 +1096,7 @@ class WanModel(torch.nn.Module):
self.use_non_blocking = False
self.video_attention_split_steps = []
self.lora_scheduling_enabled = False
# embeddings
self.patch_embedding = nn.Conv3d(
@@ -1384,7 +1385,12 @@ class WanModel(torch.nn.Module):
Returns:
List[Tensor]:
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
"""
"""
if self.lora_scheduling_enabled:
for name, submodule in self.named_modules():
if isinstance(submodule, nn.Linear):
if hasattr(submodule, 'step'):
submodule.step = current_step
# params
device = self.patch_embedding.weight.device
if freqs is not None and freqs.device != device: