From fccb2375424b4df514eb31fa94f13d415366780f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 8 Oct 2023 16:01:53 -0400 Subject: [PATCH] fix lora unload --- .../Lora/extra_networks_lora.py | 24 ++++++++++++++----- extensions-builtin/Lora/lora_patches.py | 19 +++++++++++++++ extensions-builtin/Lora/networks.py | 4 ++-- modules/processing.py | 2 +- 4 files changed, 40 insertions(+), 9 deletions(-) diff --git a/extensions-builtin/Lora/extra_networks_lora.py b/extensions-builtin/Lora/extra_networks_lora.py index a80d97b54..c7bc14d03 100644 --- a/extensions-builtin/Lora/extra_networks_lora.py +++ b/extensions-builtin/Lora/extra_networks_lora.py @@ -5,9 +5,13 @@ from modules import extra_networks, shared class ExtraNetworkLora(extra_networks.ExtraNetwork): + def __init__(self): super().__init__('lora') + self.active = False self.errors = {} + networks.originals = lora_patches.LoraPatches() + """mapping of network names to the number of errors the network had during operation""" def activate(self, p, params_list): @@ -18,7 +22,10 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): p.all_prompts = [x + f"" for x in p.all_prompts] params_list.append(extra_networks.ExtraNetworkParams(items=[additional, shared.opts.extra_networks_default_multiplier])) if len(params_list) > 0: - networks.originals = lora_patches.LoraPatches() + self.active = True + if networks.debug: + shared.log.debug("LoRA activate") + networks.originals.apply() names = [] te_multipliers = [] unet_multipliers = [] @@ -51,13 +58,18 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): network_hashes.append(f"{alias}: {shorthash}") if network_hashes: p.extra_generation_params["Lora hashes"] = ", ".join(network_hashes) - shared.log.info(f'Applying LoRA: {names} patch={t1-t0:.2f}s load={t2-t1:.2f}s') - + if len(names) > 0: + shared.log.info(f'Applying LoRA: {names} patch={t1-t0:.2f}s load={t2-t1:.2f}s') + elif self.active: + self.active = False def deactivate(self, p): - networks.originals.undo() - if networks.debug: - shared.log.debug(f"LoRA timers: load={networks.timer['load']:.2f}s apply={networks.timer['apply']:.2f}s restore={networks.timer['restore']:.2f}s") + if not self.active and getattr(networks, "originals", None ) is not None: + networks.originals.undo() + if networks.debug: + shared.log.debug("LoRA deactivate") + if self.active and networks.debug: + shared.log.debug(f"LoRA end: load={networks.timer['load']:.2f}s apply={networks.timer['apply']:.2f}s restore={networks.timer['restore']:.2f}s") if self.errors: p.comment("Networks with errors: " + ", ".join(f"{k} ({v})" for k, v in self.errors.items())) for k, v in self.errors.items(): diff --git a/extensions-builtin/Lora/lora_patches.py b/extensions-builtin/Lora/lora_patches.py index 14d8a18ac..7eba99e18 100644 --- a/extensions-builtin/Lora/lora_patches.py +++ b/extensions-builtin/Lora/lora_patches.py @@ -5,6 +5,21 @@ from modules import patches class LoraPatches: def __init__(self): + self.active = False + self.Linear_forward = None + self.Linear_load_state_dict = None + self.Conv2d_forward = None + self.Conv2d_load_state_dict = None + self.GroupNorm_forward = None + self.GroupNorm_load_state_dict = None + self.LayerNorm_forward = None + self.LayerNorm_load_state_dict = None + self.MultiheadAttention_forward = None + self.MultiheadAttention_load_state_dict = None + + def apply(self): + if self.active: + return self.Linear_forward = patches.patch(__name__, torch.nn.Linear, 'forward', networks.network_Linear_forward) self.Linear_load_state_dict = patches.patch(__name__, torch.nn.Linear, '_load_from_state_dict', networks.network_Linear_load_state_dict) self.Conv2d_forward = patches.patch(__name__, torch.nn.Conv2d, 'forward', networks.network_Conv2d_forward) @@ -18,8 +33,11 @@ class LoraPatches: networks.timer['load'] = 0 networks.timer['apply'] = 0 networks.timer['restore'] = 0 + self.active = True def undo(self): + if not self.active: + return self.Linear_forward = patches.undo(__name__, torch.nn.Linear, 'forward') # pylint: disable=E1128 self.Linear_load_state_dict = patches.undo(__name__, torch.nn.Linear, '_load_from_state_dict') # pylint: disable=E1128 self.Conv2d_forward = patches.undo(__name__, torch.nn.Conv2d, 'forward') # pylint: disable=E1128 @@ -31,3 +49,4 @@ class LoraPatches: self.MultiheadAttention_forward = patches.undo(__name__, torch.nn.MultiheadAttention, 'forward') # pylint: disable=E1128 self.MultiheadAttention_load_state_dict = patches.undo(__name__, torch.nn.MultiheadAttention, '_load_from_state_dict') # pylint: disable=E1128 patches.originals.pop(__name__, None) + self.active = False diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 8e4e01e52..641cc8aac 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -173,8 +173,8 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No while len(lora_cache) > shared.opts.lora_in_memory_limit: name = next(iter(lora_cache)) lora_cache.pop(name, None) - if debug: - shared.log.debug(f'LoRA cache: {list(lora_cache)}') + if len(loaded_networks) > 0 and debug: + shared.log.debug(f'LoRA loaded={len(loaded_networks)} cache={list(lora_cache)}') devices.torch_gc() if recompile_model: diff --git a/modules/processing.py b/modules/processing.py index 453d2b033..6756b8040 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -902,7 +902,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: if shared.opts.grid_save: images.save_image(grid, p.outpath_grids, "", p.all_seeds[0], p.all_prompts[0], shared.opts.grid_format, info=infotext(-1), short_filename=not shared.opts.grid_extended_filename, p=p, grid=True, suffix="-grid") # main save grid - if not p.disable_extra_networks and extra_network_data: + if not p.disable_extra_networks: modules.extra_networks.deactivate(p, extra_network_data) res = Processed(