Restore LoRA step scheduling functionality

This commit is contained in:
kijai
2025-12-13 00:54:22 +02:00
parent 4a3ab6958a
commit 164a6bbebd
2 changed files with 7 additions and 4 deletions
+5
View File
@@ -264,6 +264,11 @@ class CustomLinear(nn.Linear):
del weight, input, bias
return out
def update_lora_step(module, step):
for name, submodule in module.named_modules():
if isinstance(submodule, CustomLinear) and hasattr(submodule, "_step"):
submodule._step.fill_(step)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
if hasattr(submodule, "lora_diffs"):
+2 -4
View File
@@ -21,6 +21,7 @@ from ...utils import log, get_module_memory_mb
from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCacheState, relative_l1_distance
from ...multitalk.multitalk import get_attn_map_with_target
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from ...custom_linear import update_lora_step
from ...MTV.mtv import apply_rotary_emb
from comfy.ldm.flux.math import apply_rope1 as apply_rope_comfy1
@@ -2251,10 +2252,7 @@ class WanModel(torch.nn.Module):
ip_scale = fantasy_portrait_input.get("strength", 1.0)
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
update_lora_step(self, current_step)
# lynx
lynx_x_ip = lynx_ref_feature = lynx_ref_buffer = lynx_ref_feature_extractor = None