From 383d7052ac135db857f2ab21e82f50437148ab07 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 9 Dec 2024 15:23:22 -0500 Subject: [PATCH] lora split te apply Signed-off-by: Vladimir Mandic --- modules/extra_networks.py | 45 ++++++++++++++++------------- modules/lora/extra_networks_lora.py | 6 ++-- modules/lora/networks.py | 19 +++++++----- modules/processing_args.py | 3 +- modules/processing_diffusers.py | 2 +- 5 files changed, 43 insertions(+), 32 deletions(-) diff --git a/modules/extra_networks.py b/modules/extra_networks.py index e96d2e5b7..fe141cca1 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -1,6 +1,7 @@ import re +import inspect from collections import defaultdict -from modules import errors, shared, devices +from modules import errors, shared extra_network_registry = {} @@ -74,7 +75,7 @@ def is_stepwise(en_obj): return any([len(str(x).split("@")) > 1 for x in all_args]) # noqa C419 # pylint: disable=use-a-generator -def activate(p, extra_network_data=None, step=0): +def activate(p, extra_network_data=None, step=0, include=[], exclude=[]): """call activate for extra networks in extra_network_data in specified order, then call activate for all remaining registered networks with an empty argument list""" if p.disable_extra_networks: return @@ -89,25 +90,29 @@ def activate(p, extra_network_data=None, step=0): shared.log.warning("Composable LoRA not compatible with 'lora_force_diffusers'") stepwise = False shared.opts.data['lora_functional'] = stepwise or functional - with devices.autocast(): - for extra_network_name, extra_network_args in extra_network_data.items(): - extra_network = extra_network_registry.get(extra_network_name, None) - if extra_network is None: - errors.log.warning(f"Skipping unknown extra network: {extra_network_name}") - continue - try: - extra_network.activate(p, extra_network_args, step=step) - except Exception as e: - errors.display(e, f"Activating network: type={extra_network_name} args:{extra_network_args}") - for extra_network_name, extra_network in extra_network_registry.items(): - args = extra_network_data.get(extra_network_name, None) - if args is not None: - continue - try: - extra_network.activate(p, []) - except Exception as e: - errors.display(e, f"Activating network: type={extra_network_name}") + for extra_network_name, extra_network_args in extra_network_data.items(): + extra_network = extra_network_registry.get(extra_network_name, None) + if extra_network is None: + errors.log.warning(f"Skipping unknown extra network: {extra_network_name}") + continue + try: + signature = list(inspect.signature(extra_network.activate).parameters) + if 'include' in signature and 'exclude' in signature: + extra_network.activate(p, extra_network_args, step=step, include=include, exclude=exclude) + else: + extra_network.activate(p, extra_network_args, step=step) + except Exception as e: + errors.display(e, f"Activating network: type={extra_network_name} args:{extra_network_args}") + + for extra_network_name, extra_network in extra_network_registry.items(): + args = extra_network_data.get(extra_network_name, None) + if args is not None: + continue + try: + extra_network.activate(p, []) + except Exception as e: + errors.display(e, f"Activating network: type={extra_network_name}") p.network_data = extra_network_data if stepwise: diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 4ce7a94a9..135df1ccb 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -112,7 +112,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): self.model = None self.errors = {} - def activate(self, p, params_list, step=0): + def activate(self, p, params_list, step=0, include=[], exclude=[]): self.errors.clear() if self.active: if self.model != shared.opts.sd_model_checkpoint: # reset if model changed @@ -123,8 +123,8 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): self.model = shared.opts.sd_model_checkpoint names, te_multipliers, unet_multipliers, dyn_dims = parse(p, params_list, step) networks.network_load(names, te_multipliers, unet_multipliers, dyn_dims) # load - networks.network_activate() - if len(networks.loaded_networks) > 0 and step == 0: + networks.network_activate(include, exclude) + if len(networks.loaded_networks) > 0 and len(networks.applied_layers) > 0 and step == 0: infotext(p) prompt(p) shared.log.info(f'Load network: type=LoRA apply={[n.name for n in networks.loaded_networks]} te={te_multipliers} unet={unet_multipliers} time={networks.timer.summary}') diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 805b24b52..edd82f3e4 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -19,6 +19,7 @@ extra_network_lora = ExtraNetworkLora() available_networks = {} available_network_aliases = {} loaded_networks: List[network.Network] = [] +applied_layers: list[str] = [] bnb = None lora_cache = {} diffuser_loaded = [] @@ -465,7 +466,7 @@ def network_deactivate(): task = None pbar = nullcontext() with devices.inference_context(), pbar: - applied_layers = [] + applied_layers.clear() weights_devices = [] weights_dtypes = [] for component in modules.keys(): @@ -498,7 +499,7 @@ def network_deactivate(): sd_models.set_diffuser_offload(sd_model, op="model") -def network_activate(): +def network_activate(include=[], exclude=[]): t0 = time.time() timer.clear(complete=True) sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility @@ -506,10 +507,12 @@ def network_activate(): sd_models.disable_offload(sd_model) sd_models.move_model(sd_model, device=devices.cpu) modules = {} - for component_name in ['text_encoder','text_encoder_2', 'unet', 'transformer']: - component = getattr(sd_model, component_name, None) + components = include if len(include) > 0 else ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer'] + components = [x for x in components if x not in exclude] + for name in components: + component = getattr(sd_model, name, None) if component is not None and hasattr(component, 'named_modules'): - modules[component_name] = list(component.named_modules()) + modules[name] = list(component.named_modules()) total = sum(len(x) for x in modules.values()) if len(loaded_networks) > 0: pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=activate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console) @@ -519,7 +522,7 @@ def network_activate(): pbar = nullcontext() with devices.inference_context(), pbar: wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) if len(loaded_networks) > 0 else () - applied_layers = [] + applied_layers.clear() backup_size = 0 weights_devices = [] weights_dtypes = [] @@ -546,10 +549,12 @@ def network_activate(): module.network_current_names = wanted_names if task is not None: pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={total} apply={len(applied_layers)} backup={backup_size}') + if task is not None and len(applied_layers) == 0: + pbar.remove_task(task) # hide progress bar for no action weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403 # pylint: disable=R1718 timer.activate = time.time() - t0 if debug and len(loaded_networks) > 0: - shared.log.debug(f'Load network: type=LoRA networks={len(loaded_networks)} modules={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') + shared.log.debug(f'Load network: type=LoRA networks={len(loaded_networks)} components={components} modules={total} apply={len(applied_layers)} device={weights_devices} dtype={weights_dtypes} backup={backup_size} fuse={shared.opts.lora_fuse_diffusers} time={timer.summary}') modules.clear() if shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/processing_args.py b/modules/processing_args.py index 93b0bf9b2..e7f53ba8e 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -6,7 +6,7 @@ import time import inspect import torch import numpy as np -from modules import shared, errors, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, prompt_parser_diffusers, timer +from modules import shared, errors, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, prompt_parser_diffusers, timer, extra_networks from modules.processing_callbacks import diffusers_callback_legacy, diffusers_callback, set_callbacks_p from modules.processing_helpers import resize_hires, fix_prompts, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, get_generator, set_latents, apply_circular # pylint: disable=unused-import from modules.api import helpers @@ -134,6 +134,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 else: prompt_parser_diffusers.embedder = None + extra_networks.activate(p, include=['text_encoder', 'text_encoder_2', 'text_encoder_3']) if 'prompt' in possible: if 'OmniGen' in model.__class__.__name__: prompts = [p.replace('|image|', '<|image_1|>') for p in prompts] diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index d22a9de97..627eb281f 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -89,7 +89,7 @@ def process_base(p: processing.StableDiffusionProcessing): sd_models.move_model(shared.sd_model.unet, devices.device) if hasattr(shared.sd_model, 'transformer'): sd_models.move_model(shared.sd_model.transformer, devices.device) - extra_networks.activate(p) + extra_networks.activate(p, exclude=['text_encoder', 'text_encoder_2']) hidiffusion.apply(p, shared.sd_model_type) # if 'image' in base_args: # base_args['image'] = set_latents(p)