Restore LoRA step scheduling functionality
This commit is contained in:
@@ -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"):
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user