Restore LoRA step scheduling functionality
This commit is contained in:
@@ -264,6 +264,11 @@ class CustomLinear(nn.Linear):
|
|||||||
del weight, input, bias
|
del weight, input, bias
|
||||||
return out
|
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):
|
def remove_lora_from_module(module):
|
||||||
for name, submodule in module.named_modules():
|
for name, submodule in module.named_modules():
|
||||||
if hasattr(submodule, "lora_diffs"):
|
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 ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCacheState, relative_l1_distance
|
||||||
from ...multitalk.multitalk import get_attn_map_with_target
|
from ...multitalk.multitalk import get_attn_map_with_target
|
||||||
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
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 ...MTV.mtv import apply_rotary_emb
|
||||||
from comfy.ldm.flux.math import apply_rope1 as apply_rope_comfy1
|
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)
|
ip_scale = fantasy_portrait_input.get("strength", 1.0)
|
||||||
|
|
||||||
if self.lora_scheduling_enabled:
|
if self.lora_scheduling_enabled:
|
||||||
for name, submodule in self.named_modules():
|
update_lora_step(self, current_step)
|
||||||
if isinstance(submodule, nn.Linear):
|
|
||||||
if hasattr(submodule, 'step'):
|
|
||||||
submodule.step = current_step
|
|
||||||
|
|
||||||
# lynx
|
# lynx
|
||||||
lynx_x_ip = lynx_ref_feature = lynx_ref_buffer = lynx_ref_feature_extractor = None
|
lynx_x_ip = lynx_ref_feature = lynx_ref_buffer = lynx_ref_feature_extractor = None
|
||||||
|
|||||||
Reference in New Issue
Block a user