diff --git a/fp8_optimization.py b/fp8_optimization.py index 6733f07..3d6c21f 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -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): diff --git a/nodes.py b/nodes.py index 2a1a096..763eed8 100644 --- a/nodes.py +++ b/nodes.py @@ -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"]) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index c1ce758..0bece85 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4300883..bd31312 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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: