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:
+7
-2
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user