diff --git a/extensions-builtin/Lora/network.py b/extensions-builtin/Lora/network.py index a6579ae90..bc60389fd 100644 --- a/extensions-builtin/Lora/network.py +++ b/extensions-builtin/Lora/network.py @@ -65,6 +65,7 @@ class Network: # LoraModule self.unet_multiplier = [1.0] * 3 self.dyn_dim = None self.modules = {} + self.bundle_embeddings = {} self.mtime = None self.mentioned_name = None """the text that was used to add the network to prompt - can be either name or an alias""" diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 71b5b29dc..4c6162677 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -1,504 +1,513 @@ -from typing import Union, List -import os -import re -import time -import concurrent -import lora_patches -import network -import network_lora -import network_hada -import network_ia3 -import network_oft -import network_lokr -import network_full -import network_norm -import network_glora -import network_overrides -import lora_convert -import torch -import diffusers.models.lora -from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, files_cache - - -debug = os.environ.get('SD_LORA_DEBUG', None) is not None -originals: lora_patches.LoraPatches = None -extra_network_lora = None -available_networks = {} -available_network_aliases = {} -loaded_networks: List[network.Network] = [] -timer = { 'load': 0, 'apply': 0, 'restore': 0 } -# networks_in_memory = {} -lora_cache = {} -available_network_hash_lookup = {} -forbidden_network_aliases = {} -re_network_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") -module_types = [ - network_lora.ModuleTypeLora(), - network_hada.ModuleTypeHada(), - network_ia3.ModuleTypeIa3(), - network_oft.ModuleTypeOFT(), - network_lokr.ModuleTypeLokr(), - network_full.ModuleTypeFull(), - network_norm.ModuleTypeNorm(), - network_glora.ModuleTypeGLora(), -] -convert_diffusers_name_to_compvis = lora_convert.convert_diffusers_name_to_compvis # supermerger compatibility item - - -def assign_network_names_to_compvis_modules(sd_model): - network_layer_mapping = {} - if shared.native: - if not hasattr(shared.sd_model, 'text_encoder') or not hasattr(shared.sd_model, 'unet'): - sd_model.network_layer_mapping = {} - return - for name, module in shared.sd_model.text_encoder.named_modules(): - prefix = "lora_te1_" if shared.sd_model_type == "sdxl" else "lora_te_" - network_name = prefix + name.replace(".", "_") - network_layer_mapping[network_name] = module - module.network_layer_name = network_name - if shared.sd_model_type == "sdxl": - for name, module in shared.sd_model.text_encoder_2.named_modules(): - network_name = "lora_te2_" + name.replace(".", "_") - network_layer_mapping[network_name] = module - module.network_layer_name = network_name - for name, module in shared.sd_model.unet.named_modules(): - network_name = "lora_unet_" + name.replace(".", "_") - network_layer_mapping[network_name] = module - module.network_layer_name = network_name - else: - if not hasattr(shared.sd_model, 'cond_stage_model'): - sd_model.network_layer_mapping = {} - return - for name, module in shared.sd_model.cond_stage_model.wrapped.named_modules(): - network_name = name.replace(".", "_") - network_layer_mapping[network_name] = module - module.network_layer_name = network_name - for name, module in shared.sd_model.model.named_modules(): - network_name = name.replace(".", "_") - network_layer_mapping[network_name] = module - module.network_layer_name = network_name - sd_model.network_layer_mapping = network_layer_mapping - - -def load_diffusers(name, network_on_disk, lora_scale=1.0) -> network.Network: - t0 = time.time() - cached = lora_cache.get(name, None) - # if debug: - shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" type=diffusers {"cached" if cached else ""} fuse={shared.opts.lora_fuse_diffusers}') - if cached is not None: - return cached - if not shared.native: - return None - if not hasattr(shared.sd_model, 'load_lora_weights'): - shared.log.error(f"LoRA load failed: class={shared.sd_model.__class__} does not implement load lora") - return None - try: - shared.sd_model.load_lora_weights(network_on_disk.filename) - except Exception as e: - errors.display(e, "LoRA") - return None - if shared.opts.lora_fuse_diffusers: - shared.sd_model.fuse_lora(lora_scale=lora_scale) - net = network.Network(name, network_on_disk) - net.mtime = os.path.getmtime(network_on_disk.filename) - lora_cache[name] = net - t1 = time.time() - timer['load'] += t1 - t0 - return net - - -def load_network(name, network_on_disk) -> network.Network: - t0 = time.time() - cached = lora_cache.get(name, None) - if debug: - shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') - if cached is not None: - return cached - net = network.Network(name, network_on_disk) - net.mtime = os.path.getmtime(network_on_disk.filename) - sd = sd_models.read_state_dict(network_on_disk.filename) - assign_network_names_to_compvis_modules(shared.sd_model) # this should not be needed but is here as an emergency fix for an unknown error people are experiencing in 1.2.0 - keys_failed_to_match = {} - matched_networks = {} - convert = lora_convert.KeyConvert() - for key_network, weight in sd.items(): - parts = key_network.split('.') - if len(parts) > 5: # messy handler for diffusers peft lora - key_network_without_network_parts = '_'.join(parts[:-2]) - if not key_network_without_network_parts.startswith('lora_'): - key_network_without_network_parts = 'lora_' + key_network_without_network_parts - network_part = '.'.join(parts[-2:]).replace('lora_A', 'lora_down').replace('lora_B', 'lora_up') - else: - key_network_without_network_parts, network_part = key_network.split(".", 1) - # if debug: - # shared.log.debug(f'LoRA load: name="{name}" full={key_network} network={network_part} key={key_network_without_network_parts}') - key, sd_module = convert(key_network_without_network_parts) # Now returns lists - if sd_module[0] is None: - keys_failed_to_match[key_network] = key - continue - for k, module in zip(key, sd_module): - if k not in matched_networks: - matched_networks[k] = network.NetworkWeights(network_key=key_network, sd_key=k, w={}, sd_module=module) - matched_networks[k].w[network_part] = weight - for key, weights in matched_networks.items(): - net_module = None - for nettype in module_types: - net_module = nettype.create_module(net, weights) - if net_module is not None: - break - if net_module is None: - shared.log.error(f'LoRA unhandled: name={name} key={key} weights={weights.w.keys()}') - else: - net.modules[key] = net_module - if len(keys_failed_to_match) > 0: - shared.log.warning(f"LoRA file={network_on_disk.filename} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}") - if debug: - shared.log.debug(f"LoRA file={network_on_disk.filename} unmatched={keys_failed_to_match}") - elif debug: - shared.log.debug(f"LoRA file={network_on_disk.filename} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}") - lora_cache[name] = net - t1 = time.time() - timer['load'] += t1 - t0 - return net - - -def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=None): - networks_on_disk = [available_network_aliases.get(name, None) for name in names] - if any(x is None for x in networks_on_disk): - list_available_networks() - networks_on_disk = [available_network_aliases.get(name, None) for name in names] - failed_to_load_networks = [] - recompile_model = False - if shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled: - if len(names) == len(shared.compiled_model_state.lora_model): - for i, name in enumerate(names): - if shared.compiled_model_state.lora_model[i] != f"{name}:{te_multipliers[i] if te_multipliers else 1.0}": - recompile_model = True - shared.compiled_model_state.lora_model = [] - break - if not recompile_model: - if len(loaded_networks) > 0 and debug: - shared.log.debug('Model Compile: Skipping LoRa loading') - return - else: - recompile_model = True - shared.compiled_model_state.lora_model = [] - if recompile_model: - backup_cuda_compile = shared.opts.cuda_compile - sd_models.unload_model_weights(op='model') - shared.opts.cuda_compile = False - sd_models.reload_model_weights(op='model') - shared.opts.cuda_compile = backup_cuda_compile - - loaded_networks.clear() - for i, (network_on_disk, name) in enumerate(zip(networks_on_disk, names)): - net = None - if network_on_disk is not None: - shorthash = getattr(network_on_disk, 'shorthash', '').lower() - if debug: - shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') - try: - if recompile_model: - shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else 1.0}") - if shared.native and shared.opts.lora_force_diffusers: # OpenVINO only works with Diffusers LoRa loading - net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else 1.0) - elif shared.native and network_overrides.check_override(shorthash): - net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else 1.0) - else: - net = load_network(name, network_on_disk) - except Exception as e: - shared.log.error(f"LoRA load failed: file={network_on_disk.filename} {e}") - if debug: - errors.display(e, f"LoRA load failed file={network_on_disk.filename}") - continue - net.mentioned_name = name - network_on_disk.read_hash() - if net is None: - failed_to_load_networks.append(name) - shared.log.error(f"LoRA unknown type: network={name}") - continue - net.te_multiplier = te_multipliers[i] if te_multipliers else 1.0 - net.unet_multiplier = unet_multipliers[i] if unet_multipliers else 1.0 - net.dyn_dim = dyn_dims[i] if dyn_dims else 1.0 - loaded_networks.append(net) - - while len(lora_cache) > shared.opts.lora_in_memory_limit: - name = next(iter(lora_cache)) - lora_cache.pop(name, None) - 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: - shared.log.info("LoRA recompiling model") - backup_lora_model = shared.compiled_model_state.lora_model - if shared.opts.cuda_compile: - shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) - - shared.compiled_model_state.lora_model = backup_lora_model - - -def network_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, torch.nn.MultiheadAttention, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv]): - t0 = time.time() - weights_backup = getattr(self, "network_weights_backup", None) - bias_backup = getattr(self, "network_bias_backup", None) - if weights_backup is None and bias_backup is None: - return - # if debug: - # shared.log.debug('LoRA restore weights') - if weights_backup is not None: - if isinstance(self, torch.nn.MultiheadAttention): - self.in_proj_weight.copy_(weights_backup[0]) - self.out_proj.weight.copy_(weights_backup[1]) - else: - self.weight.copy_(weights_backup) - if bias_backup is not None: - if isinstance(self, torch.nn.MultiheadAttention): - self.out_proj.bias.copy_(bias_backup) - else: - self.bias.copy_(bias_backup) - else: - if isinstance(self, torch.nn.MultiheadAttention): - self.out_proj.bias = None - else: - self.bias = None - t1 = time.time() - timer['restore'] += t1 - t0 - - -def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, torch.nn.MultiheadAttention, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv]): - """ - Applies the currently selected set of networks to the weights of torch layer self. - If weights already have this particular set of networks applied, does nothing. - If not, restores orginal weights from backup and alters weights according to networks. - """ - network_layer_name = getattr(self, 'network_layer_name', None) - if network_layer_name is None: - return - t0 = time.time() - current_names = getattr(self, "network_current_names", ()) - wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) - weights_backup = getattr(self, "network_weights_backup", None) - if weights_backup is None and wanted_names != (): # pylint: disable=C1803 - if current_names != (): - raise RuntimeError("no backup weights found and current weights are not unchanged") - if isinstance(self, torch.nn.MultiheadAttention): - weights_backup = (self.in_proj_weight.to(devices.cpu, copy=True), self.out_proj.weight.to(devices.cpu, copy=True)) - else: - weights_backup = self.weight.to(devices.cpu, copy=True) - self.network_weights_backup = weights_backup - bias_backup = getattr(self, "network_bias_backup", None) - if bias_backup is None: - if isinstance(self, torch.nn.MultiheadAttention) and self.out_proj.bias is not None: - bias_backup = self.out_proj.bias.to(devices.cpu, copy=True) - elif getattr(self, 'bias', None) is not None: - bias_backup = self.bias.to(devices.cpu, copy=True) - else: - bias_backup = None - self.network_bias_backup = bias_backup - - if current_names != wanted_names: - network_restore_weights_from_backup(self) - for net in loaded_networks: - # default workflow where module is known and has weights - module = net.modules.get(network_layer_name, None) - if module is not None and hasattr(self, 'weight'): - try: - with devices.inference_context(): - updown, ex_bias = module.calc_updown(self.weight) - if len(self.weight.shape) == 4 and self.weight.shape[1] == 9: - # inpainting model. zero pad updown to make channel[1] 4 to 9 - updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable - self.weight = torch.nn.Parameter(self.weight + updown) - if ex_bias is not None and hasattr(self, 'bias'): - if self.bias is None: - self.bias = torch.nn.Parameter(ex_bias) - else: - self.bias += ex_bias - except RuntimeError as e: - extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 - if debug: - module_name = net.modules.get(network_layer_name, None) - shared.log.error(f"LoRA apply weight name={net.name} module={module_name} layer={network_layer_name} {e}") - errors.display(e, 'LoRA apply weight') - raise RuntimeError('LoRA apply weight') from e - continue - # alternative workflow looking at _*_proj layers - module_q = net.modules.get(network_layer_name + "_q_proj", None) - module_k = net.modules.get(network_layer_name + "_k_proj", None) - module_v = net.modules.get(network_layer_name + "_v_proj", None) - module_out = net.modules.get(network_layer_name + "_out_proj", None) - if isinstance(self, torch.nn.MultiheadAttention) and module_q and module_k and module_v and module_out: - try: - with devices.inference_context(): - updown_q, _ = module_q.calc_updown(self.in_proj_weight) - updown_k, _ = module_k.calc_updown(self.in_proj_weight) - updown_v, _ = module_v.calc_updown(self.in_proj_weight) - updown_qkv = torch.vstack([updown_q, updown_k, updown_v]) - updown_out, ex_bias = module_out.calc_updown(self.out_proj.weight) - self.in_proj_weight += updown_qkv - self.out_proj.weight += updown_out - if ex_bias is not None: - if self.out_proj.bias is None: - self.out_proj.bias = torch.nn.Parameter(ex_bias) - else: - self.out_proj.bias += ex_bias - except RuntimeError as e: - if debug: - shared.log.debug(f"LoRA network={net.name} layer={network_layer_name} {e}") - extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 - continue - if module is None: - continue - shared.log.warning(f"LoRA network={net.name} layer={network_layer_name} unsupported operation") - extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 - self.network_current_names = wanted_names - t1 = time.time() - timer['apply'] += t1 - t0 - - -def network_forward(module, input, original_forward): # pylint: disable=W0622 - """ - Old way of applying Lora by executing operations during layer's forward. - Stacking many loras this way results in big performance degradation. - """ - if len(loaded_networks) == 0: - return original_forward(module, input) - input = devices.cond_cast_unet(input) - network_restore_weights_from_backup(module) - network_reset_cached_weight(module) - y = original_forward(module, input) - network_layer_name = getattr(module, 'network_layer_name', None) - for lora in loaded_networks: - module = lora.modules.get(network_layer_name, None) - if module is None: - continue - y = module.forward(input, y) - return y - - -def network_reset_cached_weight(self: Union[torch.nn.Conv2d, torch.nn.Linear]): - self.network_current_names = () - self.network_weights_backup = None - - -def network_Linear_forward(self, input): # pylint: disable=W0622 - if shared.opts.lora_functional: - return network_forward(self, input, originals.Linear_forward) - network_apply_weights(self) - return originals.Linear_forward(self, input) - - -def network_Linear_load_state_dict(self, *args, **kwargs): - network_reset_cached_weight(self) - return originals.Linear_load_state_dict(self, *args, **kwargs) - - -def network_Conv2d_forward(self, input): # pylint: disable=W0622 - if shared.opts.lora_functional: - return network_forward(self, input, originals.Conv2d_forward) - network_apply_weights(self) - return originals.Conv2d_forward(self, input) - - -def network_Conv2d_load_state_dict(self, *args, **kwargs): - network_reset_cached_weight(self) - return originals.Conv2d_load_state_dict(self, *args, **kwargs) - - -def network_GroupNorm_forward(self, input): # pylint: disable=W0622 - if shared.opts.lora_functional: - return network_forward(self, input, originals.GroupNorm_forward) - network_apply_weights(self) - return originals.GroupNorm_forward(self, input) - - -def network_GroupNorm_load_state_dict(self, *args, **kwargs): - network_reset_cached_weight(self) - return originals.GroupNorm_load_state_dict(self, *args, **kwargs) - - -def network_LayerNorm_forward(self, input): # pylint: disable=W0622 - if shared.opts.lora_functional: - return network_forward(self, input, originals.LayerNorm_forward) - network_apply_weights(self) - return originals.LayerNorm_forward(self, input) - - -def network_LayerNorm_load_state_dict(self, *args, **kwargs): - network_reset_cached_weight(self) - return originals.LayerNorm_load_state_dict(self, *args, **kwargs) - - -def network_MultiheadAttention_forward(self, *args, **kwargs): - network_apply_weights(self) - return originals.MultiheadAttention_forward(self, *args, **kwargs) - - -def network_MultiheadAttention_load_state_dict(self, *args, **kwargs): - network_reset_cached_weight(self) - return originals.MultiheadAttention_load_state_dict(self, *args, **kwargs) - - -def list_available_networks(): - available_networks.clear() - available_network_aliases.clear() - forbidden_network_aliases.clear() - available_network_hash_lookup.clear() - forbidden_network_aliases.update({"none": 1, "Addams": 1}) - directories = [] - if os.path.exists(shared.cmd_opts.lora_dir): - directories.append(shared.cmd_opts.lora_dir) - else: - shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') - if os.path.exists(shared.cmd_opts.lyco_dir) and shared.cmd_opts.lyco_dir != shared.cmd_opts.lora_dir: - directories.append(shared.cmd_opts.lyco_dir) - - def add_network(filename): - if not os.path.isfile(filename): - return - name = os.path.splitext(os.path.basename(filename))[0] - try: - entry = network.NetworkOnDisk(name, filename) - available_networks[entry.name] = entry - if entry.alias in available_network_aliases: - forbidden_network_aliases[entry.alias.lower()] = 1 - if shared.opts.lora_preferred_name == 'filename': - available_network_aliases[entry.name] = entry - else: - available_network_aliases[entry.alias] = entry - if entry.shorthash: - available_network_hash_lookup[entry.shorthash] = entry - except OSError as e: # should catch FileNotFoundError and PermissionError etc. - shared.log.error(f"Failed to load network {name} from {filename} {e}") - - candidates = list(files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"])) - with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - for fn in candidates: - executor.submit(add_network, fn) - shared.log.info(f'LoRA networks: available={len(available_networks)} folders={len(forbidden_network_aliases)}') - - -def infotext_pasted(infotext, params): # pylint: disable=W0613 - if "AddNet Module 1" in [x[1] for x in scripts.scripts_txt2img.infotext_fields]: - return # if the other extension is active, it will handle those fields, no need to do anything - added = [] - for k in params: - if not k.startswith("AddNet Model "): - continue - num = k[13:] - if params.get("AddNet Module " + num) != "LoRA": - continue - name = params.get("AddNet Model " + num) - if name is None: - continue - m = re_network_name.match(name) - if m: - name = m.group(1) - multiplier = params.get("AddNet Weight A " + num, "1.0") - added.append(f"") - if added: - params["Prompt"] += "\n" + "".join(added) - - -list_available_networks() +from typing import Union, List +import os +import re +import time +import concurrent +import lora_patches +import network +import network_lora +import network_hada +import network_ia3 +import network_oft +import network_lokr +import network_full +import network_norm +import network_glora +import network_overrides +import lora_convert +import torch +import diffusers.models.lora +from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, files_cache + + +debug = os.environ.get('SD_LORA_DEBUG', None) is not None +originals: lora_patches.LoraPatches = None +extra_network_lora = None +available_networks = {} +available_network_aliases = {} +loaded_networks: List[network.Network] = [] +timer = { 'load': 0, 'apply': 0, 'restore': 0 } +# networks_in_memory = {} +lora_cache = {} +available_network_hash_lookup = {} +forbidden_network_aliases = {} +re_network_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") +module_types = [ + network_lora.ModuleTypeLora(), + network_hada.ModuleTypeHada(), + network_ia3.ModuleTypeIa3(), + network_oft.ModuleTypeOFT(), + network_lokr.ModuleTypeLokr(), + network_full.ModuleTypeFull(), + network_norm.ModuleTypeNorm(), + network_glora.ModuleTypeGLora(), +] +convert_diffusers_name_to_compvis = lora_convert.convert_diffusers_name_to_compvis # supermerger compatibility item + + +def assign_network_names_to_compvis_modules(sd_model): + network_layer_mapping = {} + if shared.native: + if not hasattr(shared.sd_model, 'text_encoder') or not hasattr(shared.sd_model, 'unet'): + sd_model.network_layer_mapping = {} + return + for name, module in shared.sd_model.text_encoder.named_modules(): + prefix = "lora_te1_" if shared.sd_model_type == "sdxl" else "lora_te_" + network_name = prefix + name.replace(".", "_") + network_layer_mapping[network_name] = module + module.network_layer_name = network_name + if shared.sd_model_type == "sdxl": + for name, module in shared.sd_model.text_encoder_2.named_modules(): + network_name = "lora_te2_" + name.replace(".", "_") + network_layer_mapping[network_name] = module + module.network_layer_name = network_name + for name, module in shared.sd_model.unet.named_modules(): + network_name = "lora_unet_" + name.replace(".", "_") + network_layer_mapping[network_name] = module + module.network_layer_name = network_name + else: + if not hasattr(shared.sd_model, 'cond_stage_model'): + sd_model.network_layer_mapping = {} + return + for name, module in shared.sd_model.cond_stage_model.wrapped.named_modules(): + network_name = name.replace(".", "_") + network_layer_mapping[network_name] = module + module.network_layer_name = network_name + for name, module in shared.sd_model.model.named_modules(): + network_name = name.replace(".", "_") + network_layer_mapping[network_name] = module + module.network_layer_name = network_name + sd_model.network_layer_mapping = network_layer_mapping + + +def load_diffusers(name, network_on_disk, lora_scale=1.0) -> network.Network: + t0 = time.time() + cached = lora_cache.get(name, None) + # if debug: + shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" type=diffusers {"cached" if cached else ""} fuse={shared.opts.lora_fuse_diffusers}') + if cached is not None: + return cached + if not shared.native: + return None + if not hasattr(shared.sd_model, 'load_lora_weights'): + shared.log.error(f"LoRA load failed: class={shared.sd_model.__class__} does not implement load lora") + return None + try: + shared.sd_model.load_lora_weights(network_on_disk.filename) + except Exception as e: + errors.display(e, "LoRA") + return None + if shared.opts.lora_fuse_diffusers: + shared.sd_model.fuse_lora(lora_scale=lora_scale) + net = network.Network(name, network_on_disk) + net.mtime = os.path.getmtime(network_on_disk.filename) + lora_cache[name] = net + t1 = time.time() + timer['load'] += t1 - t0 + return net + + +def load_network(name, network_on_disk) -> network.Network: + t0 = time.time() + cached = lora_cache.get(name, None) + if debug: + shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" type=lora {"cached" if cached else ""}') + if cached is not None: + return cached + net = network.Network(name, network_on_disk) + net.mtime = os.path.getmtime(network_on_disk.filename) + sd = sd_models.read_state_dict(network_on_disk.filename) + assign_network_names_to_compvis_modules(shared.sd_model) # this should not be needed but is here as an emergency fix for an unknown error people are experiencing in 1.2.0 + keys_failed_to_match = {} + matched_networks = {} + bundle_embeddings = {} + convert = lora_convert.KeyConvert() + for key_network, weight in sd.items(): + parts = key_network.split('.') + if parts[0] == "bundle_emb": + emb_name, vec_name = parts[1], key_network.split(".", 2)[-1] + emb_dict = bundle_embeddings.get(emb_name, {}) + emb_dict[vec_name] = weight + bundle_embeddings[emb_name] = emb_dict + if len(parts) > 5: # messy handler for diffusers peft lora + key_network_without_network_parts = '_'.join(parts[:-2]) + if not key_network_without_network_parts.startswith('lora_'): + key_network_without_network_parts = 'lora_' + key_network_without_network_parts + network_part = '.'.join(parts[-2:]).replace('lora_A', 'lora_down').replace('lora_B', 'lora_up') + else: + key_network_without_network_parts, network_part = key_network.split(".", 1) + # if debug: + # shared.log.debug(f'LoRA load: name="{name}" full={key_network} network={network_part} key={key_network_without_network_parts}') + key, sd_module = convert(key_network_without_network_parts) # Now returns lists + if sd_module[0] is None: + if "bundle_emb" not in key_network: + keys_failed_to_match[key_network] = key + continue + for k, module in zip(key, sd_module): + if k not in matched_networks: + matched_networks[k] = network.NetworkWeights(network_key=key_network, sd_key=k, w={}, sd_module=module) + matched_networks[k].w[network_part] = weight + for key, weights in matched_networks.items(): + net_module = None + for nettype in module_types: + net_module = nettype.create_module(net, weights) + if net_module is not None: + break + if net_module is None: + shared.log.error(f'LoRA unhandled: name={name} key={key} weights={weights.w.keys()}') + else: + net.modules[key] = net_module + if len(keys_failed_to_match) > 0: + shared.log.warning(f"LoRA file={network_on_disk.filename} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}") + if debug: + shared.log.debug(f"LoRA file={network_on_disk.filename} unmatched={keys_failed_to_match}") + elif debug: + shared.log.debug(f"LoRA file={network_on_disk.filename} unmatched={len(keys_failed_to_match)} matched={len(matched_networks)}") + lora_cache[name] = net + t1 = time.time() + net.bundle_embeddings = bundle_embeddings + timer['load'] += t1 - t0 + return net + + +def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=None): + networks_on_disk = [available_network_aliases.get(name, None) for name in names] + if any(x is None for x in networks_on_disk): + list_available_networks() + networks_on_disk = [available_network_aliases.get(name, None) for name in names] + failed_to_load_networks = [] + recompile_model = False + if shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled: + if len(names) == len(shared.compiled_model_state.lora_model): + for i, name in enumerate(names): + if shared.compiled_model_state.lora_model[i] != f"{name}:{te_multipliers[i] if te_multipliers else 1.0}": + recompile_model = True + shared.compiled_model_state.lora_model = [] + break + if not recompile_model: + if len(loaded_networks) > 0 and debug: + shared.log.debug('Model Compile: Skipping LoRa loading') + return + else: + recompile_model = True + shared.compiled_model_state.lora_model = [] + if recompile_model: + backup_cuda_compile = shared.opts.cuda_compile + sd_models.unload_model_weights(op='model') + shared.opts.cuda_compile = False + sd_models.reload_model_weights(op='model') + shared.opts.cuda_compile = backup_cuda_compile + + loaded_networks.clear() + for i, (network_on_disk, name) in enumerate(zip(networks_on_disk, names)): + net = None + if network_on_disk is not None: + shorthash = getattr(network_on_disk, 'shorthash', '').lower() + if debug: + shared.log.debug(f'LoRA load: name="{name}" file="{network_on_disk.filename}" hash="{shorthash}"') + try: + if recompile_model: + shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else 1.0}") + if shared.native and shared.opts.lora_force_diffusers: # OpenVINO only works with Diffusers LoRa loading + net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else 1.0) + elif shared.native and network_overrides.check_override(shorthash): + net = load_diffusers(name, network_on_disk, lora_scale=te_multipliers[i] if te_multipliers else 1.0) + else: + net = load_network(name, network_on_disk) + except Exception as e: + shared.log.error(f"LoRA load failed: file={network_on_disk.filename} {e}") + if debug: + errors.display(e, f"LoRA load failed file={network_on_disk.filename}") + continue + net.mentioned_name = name + network_on_disk.read_hash() + if net is None: + failed_to_load_networks.append(name) + shared.log.error(f"LoRA unknown type: network={name}") + continue + shared.sd_model.embedding_db.load_diffusers_embedding(None, net.bundle_embeddings) + net.te_multiplier = te_multipliers[i] if te_multipliers else 1.0 + net.unet_multiplier = unet_multipliers[i] if unet_multipliers else 1.0 + net.dyn_dim = dyn_dims[i] if dyn_dims else 1.0 + loaded_networks.append(net) + + while len(lora_cache) > shared.opts.lora_in_memory_limit: + name = next(iter(lora_cache)) + lora_cache.pop(name, None) + 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: + shared.log.info("LoRA recompiling model") + backup_lora_model = shared.compiled_model_state.lora_model + if shared.opts.cuda_compile: + shared.sd_model = sd_models_compile.compile_diffusers(shared.sd_model) + + shared.compiled_model_state.lora_model = backup_lora_model + + +def network_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, torch.nn.MultiheadAttention, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv]): + t0 = time.time() + weights_backup = getattr(self, "network_weights_backup", None) + bias_backup = getattr(self, "network_bias_backup", None) + if weights_backup is None and bias_backup is None: + return + # if debug: + # shared.log.debug('LoRA restore weights') + if weights_backup is not None: + if isinstance(self, torch.nn.MultiheadAttention): + self.in_proj_weight.copy_(weights_backup[0]) + self.out_proj.weight.copy_(weights_backup[1]) + else: + self.weight.copy_(weights_backup) + if bias_backup is not None: + if isinstance(self, torch.nn.MultiheadAttention): + self.out_proj.bias.copy_(bias_backup) + else: + self.bias.copy_(bias_backup) + else: + if isinstance(self, torch.nn.MultiheadAttention): + self.out_proj.bias = None + else: + self.bias = None + t1 = time.time() + timer['restore'] += t1 - t0 + + +def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, torch.nn.MultiheadAttention, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv]): + """ + Applies the currently selected set of networks to the weights of torch layer self. + If weights already have this particular set of networks applied, does nothing. + If not, restores orginal weights from backup and alters weights according to networks. + """ + network_layer_name = getattr(self, 'network_layer_name', None) + if network_layer_name is None: + return + t0 = time.time() + current_names = getattr(self, "network_current_names", ()) + wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in loaded_networks) + weights_backup = getattr(self, "network_weights_backup", None) + if weights_backup is None and wanted_names != (): # pylint: disable=C1803 + if current_names != (): + raise RuntimeError("no backup weights found and current weights are not unchanged") + if isinstance(self, torch.nn.MultiheadAttention): + weights_backup = (self.in_proj_weight.to(devices.cpu, copy=True), self.out_proj.weight.to(devices.cpu, copy=True)) + else: + weights_backup = self.weight.to(devices.cpu, copy=True) + self.network_weights_backup = weights_backup + bias_backup = getattr(self, "network_bias_backup", None) + if bias_backup is None: + if isinstance(self, torch.nn.MultiheadAttention) and self.out_proj.bias is not None: + bias_backup = self.out_proj.bias.to(devices.cpu, copy=True) + elif getattr(self, 'bias', None) is not None: + bias_backup = self.bias.to(devices.cpu, copy=True) + else: + bias_backup = None + self.network_bias_backup = bias_backup + + if current_names != wanted_names: + network_restore_weights_from_backup(self) + for net in loaded_networks: + # default workflow where module is known and has weights + module = net.modules.get(network_layer_name, None) + if module is not None and hasattr(self, 'weight'): + try: + with devices.inference_context(): + updown, ex_bias = module.calc_updown(self.weight) + if len(self.weight.shape) == 4 and self.weight.shape[1] == 9: + # inpainting model. zero pad updown to make channel[1] 4 to 9 + updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable + self.weight = torch.nn.Parameter(self.weight + updown) + if ex_bias is not None and hasattr(self, 'bias'): + if self.bias is None: + self.bias = torch.nn.Parameter(ex_bias) + else: + self.bias += ex_bias + except RuntimeError as e: + extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 + if debug: + module_name = net.modules.get(network_layer_name, None) + shared.log.error(f"LoRA apply weight name={net.name} module={module_name} layer={network_layer_name} {e}") + errors.display(e, 'LoRA apply weight') + raise RuntimeError('LoRA apply weight') from e + continue + # alternative workflow looking at _*_proj layers + module_q = net.modules.get(network_layer_name + "_q_proj", None) + module_k = net.modules.get(network_layer_name + "_k_proj", None) + module_v = net.modules.get(network_layer_name + "_v_proj", None) + module_out = net.modules.get(network_layer_name + "_out_proj", None) + if isinstance(self, torch.nn.MultiheadAttention) and module_q and module_k and module_v and module_out: + try: + with devices.inference_context(): + updown_q, _ = module_q.calc_updown(self.in_proj_weight) + updown_k, _ = module_k.calc_updown(self.in_proj_weight) + updown_v, _ = module_v.calc_updown(self.in_proj_weight) + updown_qkv = torch.vstack([updown_q, updown_k, updown_v]) + updown_out, ex_bias = module_out.calc_updown(self.out_proj.weight) + self.in_proj_weight += updown_qkv + self.out_proj.weight += updown_out + if ex_bias is not None: + if self.out_proj.bias is None: + self.out_proj.bias = torch.nn.Parameter(ex_bias) + else: + self.out_proj.bias += ex_bias + except RuntimeError as e: + if debug: + shared.log.debug(f"LoRA network={net.name} layer={network_layer_name} {e}") + extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 + continue + if module is None: + continue + shared.log.warning(f"LoRA network={net.name} layer={network_layer_name} unsupported operation") + extra_network_lora.errors[net.name] = extra_network_lora.errors.get(net.name, 0) + 1 + self.network_current_names = wanted_names + t1 = time.time() + timer['apply'] += t1 - t0 + + +def network_forward(module, input, original_forward): # pylint: disable=W0622 + """ + Old way of applying Lora by executing operations during layer's forward. + Stacking many loras this way results in big performance degradation. + """ + if len(loaded_networks) == 0: + return original_forward(module, input) + input = devices.cond_cast_unet(input) + network_restore_weights_from_backup(module) + network_reset_cached_weight(module) + y = original_forward(module, input) + network_layer_name = getattr(module, 'network_layer_name', None) + for lora in loaded_networks: + module = lora.modules.get(network_layer_name, None) + if module is None: + continue + y = module.forward(input, y) + return y + + +def network_reset_cached_weight(self: Union[torch.nn.Conv2d, torch.nn.Linear]): + self.network_current_names = () + self.network_weights_backup = None + + +def network_Linear_forward(self, input): # pylint: disable=W0622 + if shared.opts.lora_functional: + return network_forward(self, input, originals.Linear_forward) + network_apply_weights(self) + return originals.Linear_forward(self, input) + + +def network_Linear_load_state_dict(self, *args, **kwargs): + network_reset_cached_weight(self) + return originals.Linear_load_state_dict(self, *args, **kwargs) + + +def network_Conv2d_forward(self, input): # pylint: disable=W0622 + if shared.opts.lora_functional: + return network_forward(self, input, originals.Conv2d_forward) + network_apply_weights(self) + return originals.Conv2d_forward(self, input) + + +def network_Conv2d_load_state_dict(self, *args, **kwargs): + network_reset_cached_weight(self) + return originals.Conv2d_load_state_dict(self, *args, **kwargs) + + +def network_GroupNorm_forward(self, input): # pylint: disable=W0622 + if shared.opts.lora_functional: + return network_forward(self, input, originals.GroupNorm_forward) + network_apply_weights(self) + return originals.GroupNorm_forward(self, input) + + +def network_GroupNorm_load_state_dict(self, *args, **kwargs): + network_reset_cached_weight(self) + return originals.GroupNorm_load_state_dict(self, *args, **kwargs) + + +def network_LayerNorm_forward(self, input): # pylint: disable=W0622 + if shared.opts.lora_functional: + return network_forward(self, input, originals.LayerNorm_forward) + network_apply_weights(self) + return originals.LayerNorm_forward(self, input) + + +def network_LayerNorm_load_state_dict(self, *args, **kwargs): + network_reset_cached_weight(self) + return originals.LayerNorm_load_state_dict(self, *args, **kwargs) + + +def network_MultiheadAttention_forward(self, *args, **kwargs): + network_apply_weights(self) + return originals.MultiheadAttention_forward(self, *args, **kwargs) + + +def network_MultiheadAttention_load_state_dict(self, *args, **kwargs): + network_reset_cached_weight(self) + return originals.MultiheadAttention_load_state_dict(self, *args, **kwargs) + + +def list_available_networks(): + available_networks.clear() + available_network_aliases.clear() + forbidden_network_aliases.clear() + available_network_hash_lookup.clear() + forbidden_network_aliases.update({"none": 1, "Addams": 1}) + directories = [] + if os.path.exists(shared.cmd_opts.lora_dir): + directories.append(shared.cmd_opts.lora_dir) + else: + shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') + if os.path.exists(shared.cmd_opts.lyco_dir) and shared.cmd_opts.lyco_dir != shared.cmd_opts.lora_dir: + directories.append(shared.cmd_opts.lyco_dir) + + def add_network(filename): + if not os.path.isfile(filename): + return + name = os.path.splitext(os.path.basename(filename))[0] + try: + entry = network.NetworkOnDisk(name, filename) + available_networks[entry.name] = entry + if entry.alias in available_network_aliases: + forbidden_network_aliases[entry.alias.lower()] = 1 + if shared.opts.lora_preferred_name == 'filename': + available_network_aliases[entry.name] = entry + else: + available_network_aliases[entry.alias] = entry + if entry.shorthash: + available_network_hash_lookup[entry.shorthash] = entry + except OSError as e: # should catch FileNotFoundError and PermissionError etc. + shared.log.error(f"Failed to load network {name} from {filename} {e}") + + candidates = list(files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"])) + with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: + for fn in candidates: + executor.submit(add_network, fn) + shared.log.info(f'LoRA networks: available={len(available_networks)} folders={len(forbidden_network_aliases)}') + + +def infotext_pasted(infotext, params): # pylint: disable=W0613 + if "AddNet Module 1" in [x[1] for x in scripts.scripts_txt2img.infotext_fields]: + return # if the other extension is active, it will handle those fields, no need to do anything + added = [] + for k in params: + if not k.startswith("AddNet Model "): + continue + num = k[13:] + if params.get("AddNet Module " + num) != "LoRA": + continue + name = params.get("AddNet Model " + num) + if name is None: + continue + m = re_network_name.match(name) + if m: + name = m.group(1) + multiplier = params.get("AddNet Weight A " + num, "1.0") + added.append(f"") + if added: + params["Prompt"] += "\n" + "".join(added) + + +list_available_networks() diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 1d1e5057e..292699f2d 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -1,7 +1,6 @@ from typing import List, Union import os import time -from collections import namedtuple import torch import safetensors.torch from PIL import Image @@ -12,8 +11,6 @@ from modules.files_cache import directory_files, directory_mtime, extension_filt debug = shared.log.trace if os.environ.get('SD_TI_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: TEXTUAL INVERSION') -TokenToAdd = namedtuple("TokenToAdd", ["clip_l", "clip_g"]) - def list_embeddings(*dirs): is_ext = extension_filter(['.SAFETENSORS', '.PT' ] + ( ['.PNG', '.WEBP', '.JXL', '.AVIF', '.BIN' ] if not shared.native else [] )) @@ -21,6 +18,138 @@ def list_embeddings(*dirs): return list(filter(lambda fp: is_ext(fp) and is_not_preview(fp) and os.stat(fp).st_size > 0, directory_files(*dirs))) +def open_embeddings(filename): + """ + Load Embedding files from drive. Image embeddings not currently supported. + """ + if filename is None: + return + filenames = list(filename) + exts = [".SAFETENSORS", '.BIN', '.PT'] + embeddings = [] + skipped = [] + for _filename in filenames: + # debug(f'Embedding check: {filename}') + fullname = _filename + _filename = os.path.basename(fullname) + fn, ext = os.path.splitext(_filename) + name = os.path.basename(fn) + embedding = Embedding(vec=[], name=name, filename=fullname) + try: + if ext.upper() not in exts: + debug(f'extension `{ext}` is invalid, expected one of: {exts}') + skipped.append(name) + continue + if ext.upper() in ['.SAFETENSORS']: + with safetensors.torch.safe_open(embedding.filename, framework="pt") as f: # type: ignore + for k in f.keys(): + embedding.vec.append(f.get_tensor(k)) + else: # fallback for sd1.5 pt embeddings + vectors = torch.load(fullname, map_location=devices.device)["string_to_param"]["*"] + embedding.vec.append(vectors) + embedding.tokens = [embedding.name if i == 0 else f"{embedding.name}_{i}" for i in range(len(embedding.vec[0]))] + except: + debug(f"Could not load embedding file {fullname}") + if embedding.vec: + embeddings.append(embedding) + else: + skipped.append(name) + return embeddings, skipped + + +def convert_bundled(data): + """ + Bundled embeddings are passed as a dict from lora loading, convert to Embedding objects and pass back as list. + """ + embeddings = [] + for key in data.keys(): + embedding = Embedding(vec=[], name=key, filename=None) + for vector in data[key].values(): + embedding.vec.append(vector) + embedding.tokens = [embedding.name if i == 0 else f"{embedding.name}_{i}" for i in range(len(embedding.vec[0]))] + embeddings.append(embedding) + return embeddings, [] + + +def get_text_encoders(): + """ + Select all text encoder and tokenizer pairs from known pipelines, and index them based on the dimensionality of + their embedding layers. + """ + pipe = shared.sd_model + te_names = ["text_encoder", "text_encoder_2", "text_encoder_3"] + tokenizers_names = ["tokenizer", "tokenizer_2", "tokenizer_3"] + text_encoders = [] + tokenizers = [] + hidden_sizes = [] + for te, tok in zip(te_names, tokenizers_names): + text_encoder = getattr(pipe, te, None) + if text_encoder is None: + continue + tokenizer = getattr(pipe, tok, None) + hidden_size = text_encoder.get_input_embeddings().weight.data.shape[-1] or None + if all([text_encoder, tokenizer, hidden_size]): + text_encoders.append(text_encoder) + tokenizers.append(tokenizer) + hidden_sizes.append(hidden_size) + return text_encoders, tokenizers, hidden_sizes + + +def deref_tokenizers(tokens, tokenizers): + """ + Bundled embeddings may have the same name as a seperately loaded embedding, or there may be multiple LoRA with + differing numbers of vectors. By editing the AddedToken objects, and deleting the dict keys pointing to them, + we can ensure that a smaller embedding will not get tokenized as itself, plus the remaining vectors of the previous. + """ + for tokenizer in tokenizers: + # if tokens[0] in tokenizer.get_vocab(): + # idx = tokenizer.convert_tokens_to_ids(tokens[0]) + # debug(f"deref idx: {idx}") + # tokenizer._added_tokens_decoder[idx].content = str(time.time()) + # del tokenizer._added_tokens_encoder[tokens[0]] + + if len(tokens) > 1: + last_token = tokens[-1] + suffix = int(last_token.split("_")[-1]) + newsuffix = suffix + 1 + while last_token.replace(str(suffix), str(newsuffix)) in tokenizer.get_vocab(): + idx = tokenizer.convert_tokens_to_ids(last_token.replace(str(suffix), str(newsuffix))) + debug(f"Textual inversion: deref idx={idx}") + del tokenizer._added_tokens_encoder[last_token.replace(str(suffix), str(newsuffix))] + tokenizer._added_tokens_decoder[idx].content = str(time.time()) + newsuffix += 1 + + +def insert_tokens(embeddings: list, tokenizers: list): + """ + Add all tokens to each tokenizer in the list, with one call to each. + """ + tokens = [] + for embedding in embeddings: + tokens += embedding.tokens + for tokenizer in tokenizers: + tokenizer.add_tokens(tokens) + + +def insert_vectors(embedding, tokenizers, text_encoders, hiddensizes): + """ + Insert embeddings into the input embedding layer of a list of text encoders, matched based on embedding size, + not by name. + Future warning, if another text encoder becomes available with embedding dimensions in [768,1280,4096] + this may cause collisions. + """ + for vector, size in zip(embedding.vec, embedding.vector_sizes): + idx = hiddensizes.index(size) + unk_token_id = tokenizers[idx].convert_tokens_to_ids(tokenizers[idx].unk_token) + if text_encoders[idx].get_input_embeddings().weight.data.shape[0] != len(tokenizers[idx]): + text_encoders[idx].resize_token_embeddings(len(tokenizers[idx])) + for token, v in zip(embedding.tokens, vector.unbind()): + token_id = tokenizers[idx].convert_tokens_to_ids(token) + if token_id > unk_token_id: + text_encoders[idx].get_input_embeddings().weight.data[token_id] = v + + + class Embedding: def __init__(self, vec, name, filename=None, step=None): self.vec = vec @@ -35,6 +164,7 @@ class Embedding: self.sd_checkpoint = None self.sd_checkpoint_name = None self.optimizer_state_dict = None + self.tokens = None def save(self, filename): embedding_data = { @@ -82,6 +212,10 @@ class DirWithTextualInversionEmbeddings: def convert_embedding(tensor, text_encoder, text_encoder_2): + """ + Given a tensor of shape (b, embed_dim) and two text encoders whose tokenizers match, return a tensor with + approximately mathcing meaning, or padding if the input tensor is dissimilar to any frozen text embed + """ with torch.no_grad(): vectors = [] clip_l_embeds = text_encoder.get_input_embeddings().weight.data.clone().to(device=devices.device) @@ -91,7 +225,7 @@ def convert_embedding(tensor, text_encoder, text_encoder_2): if values < 0.707: # Arbitrary similarity to cutoff, here 45 degrees indices *= 0 # Use SDXL padding vector 0 vectors.append(indices) - vectors = torch.stack(vectors) + vectors = torch.stack(vectors).to(text_encoder_2.device) output = text_encoder_2.get_input_embeddings().weight.data[vectors] return output @@ -135,123 +269,41 @@ class EmbeddingDatabase: vec = shared.sd_model.cond_stage_model.encode_embedding_init_text(",", 1) return vec.shape[1] - def load_diffusers_embedding(self, filename: Union[str, List[str]]): - _loaded_pre = len(self.word_embeddings) - embeddings_to_load = [] - loaded_embeddings = {} - skipped_embeddings = [] + def load_diffusers_embedding(self, filename: Union[str, List[str]] = None, data: dict = None): + """ + File names take precidence over bundled embeddings passed as a dict. + Bundled embeddings are automatically set to overwrite previous embeddings. + """ + overwrite = bool(data) if not shared.sd_loaded: return 0 - tokenizer = getattr(shared.sd_model, 'tokenizer', None) - tokenizer_2 = getattr(shared.sd_model, 'tokenizer_2', None) - clip_l = getattr(shared.sd_model, 'text_encoder', None) - clip_g = getattr(shared.sd_model, 'text_encoder_2', None) - if clip_g and tokenizer_2: - model_type = 'SDXL' - elif clip_l and tokenizer: - model_type = 'SD' - else: + embeddings, skipped = open_embeddings(filename) or convert_bundled(data) + if not embeddings: return 0 - filenames = list(filename) - exts = [".SAFETENSORS", '.BIN', '.PT', '.PNG', '.WEBP', '.JXL', '.AVIF'] - for _filename in filenames: - # debug(f'Embedding check: {filename}') - fullname = _filename - _filename = os.path.basename(fullname) - fn, ext = os.path.splitext(_filename) - name = os.path.basename(fn) - embedding = Embedding(vec=None, name=name, filename=fullname) - tokenizer_vocab = tokenizer.get_vocab() - try: - if ext.upper() not in exts: - raise ValueError(f'extension `{ext}` is invalid, expected one of: {exts}') - if name in tokenizer.get_vocab() or f"{name}_1" in tokenizer.get_vocab(): - loaded_embeddings[name] = embedding - debug(f'Embedding already loaded: {name}') - embeddings_to_load.append(embedding) - except Exception as e: - skipped_embeddings.append(embedding) - debug(f'Embedding skipped: "{name}" {e}') - continue - embeddings_to_load = sorted(embeddings_to_load, key=lambda e: exts.index(os.path.splitext(e.filename)[1].upper())) - - tokens_to_add = {} - for embedding in embeddings_to_load: - try: - if embedding.name in tokens_to_add or embedding.name in loaded_embeddings: - raise ValueError('duplicate token') - embeddings_dict = {} - _, ext = os.path.splitext(embedding.filename) - if ext.upper() in ['.SAFETENSORS']: - with safetensors.torch.safe_open(embedding.filename, framework="pt") as f: # type: ignore - for k in f.keys(): - embeddings_dict[k] = f.get_tensor(k) - else: # fallback for sd1.5 pt embeddings - embeddings_dict["clip_l"] = self.load_from_file(embedding.filename, embedding.filename) - if 'emb_params' in embeddings_dict and 'clip_l' not in embeddings_dict: - embeddings_dict["clip_l"] = embeddings_dict["emb_params"] - if 'clip_l' not in embeddings_dict: - raise ValueError('Invalid Embedding, dict missing required key `clip_l`') - if 'clip_g' not in embeddings_dict and model_type == "SDXL" and shared.opts.diffusers_convert_embed: - embeddings_dict["clip_g"] = convert_embedding(embeddings_dict["clip_l"], clip_l, clip_g) - if 'clip_g' in embeddings_dict: - embedding_type = 'SDXL' - else: - embedding_type = 'SD' - if embedding_type != model_type: - raise ValueError(f'Unable to load {embedding_type} Embedding "{embedding.name}" into {model_type} Model') - _tokens_to_add = {} - for i in range(len(embeddings_dict["clip_l"])): - if len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i]): - token = embedding.name if i == 0 else f"{embedding.name}_{i}" - if token in tokenizer_vocab: - raise RuntimeError(f'Multi-Vector Embedding would add pre-existing Token in Vocabulary: {token}') - if token in tokens_to_add: - raise RuntimeError(f'Multi-Vector Embedding would add duplicate Token to Add: {token}') - _tokens_to_add[token] = TokenToAdd( - embeddings_dict["clip_l"][i], - embeddings_dict["clip_g"][i] if 'clip_g' in embeddings_dict else None - ) - if not _tokens_to_add: - raise ValueError('no valid tokens to add') - tokens_to_add.update(_tokens_to_add) - loaded_embeddings[embedding.name] = embedding - except Exception as e: - debug(f"Embedding loading: {embedding.filename} {e}") - continue - if len(tokens_to_add) > 0: - tokenizer.add_tokens(list(tokens_to_add.keys())) - clip_l.resize_token_embeddings(len(tokenizer)) - if model_type == 'SDXL': - tokenizer_2.add_tokens(list(tokens_to_add.keys())) # type: ignore - clip_g.resize_token_embeddings(len(tokenizer_2)) # type: ignore - unk_token_id = tokenizer.convert_tokens_to_ids(tokenizer.unk_token) - for token, data in tokens_to_add.items(): - token_id = tokenizer.convert_tokens_to_ids(token) - if token_id > unk_token_id: - clip_l.get_input_embeddings().weight.data[token_id] = data.clip_l - if model_type == 'SDXL': - clip_g.get_input_embeddings().weight.data[token_id] = data.clip_g # type: ignore - - for embedding in loaded_embeddings.values(): - if not embedding: - continue - self.register_embedding(embedding, shared.sd_model) - if embedding in embeddings_to_load: - embeddings_to_load.remove(embedding) - skipped_embeddings.extend(embeddings_to_load) - for embedding in skipped_embeddings: - if loaded_embeddings.get(embedding.name, None) == embedding: - continue - self.skipped_embeddings[embedding.name] = embedding - try: - if model_type == 'SD': - debug(f"Embeddings loaded: text-encoder={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]}") - if model_type == 'SDXL': - debug(f"Embeddings loaded: text-encoder-1={shared.sd_model.text_encoder.get_input_embeddings().weight.data.shape[0]} text-encoder-2={shared.sd_model.text_encoder_2.get_input_embeddings().weight.data.shape[0]}") - except Exception: - pass - return len(self.word_embeddings) - _loaded_pre + text_encoders, tokenizers, hiddensizes = get_text_encoders() + if not all([text_encoders, tokenizers, hiddensizes]): + return 0 + for embedding in embeddings: + embedding.vector_sizes = [v.shape[-1] for v in embedding.vec] + if shared.opts.diffusers_convert_embed and 768 in hiddensizes and 1280 in hiddensizes and 1280 not in embedding.vector_sizes and 768 in embedding.vector_sizes: + embedding.vec.append( + convert_embedding(embedding.vec[embedding.vector_sizes.index(768)], text_encoders[hiddensizes.index(768)], + text_encoders[hiddensizes.index(1280)])) + embedding.vector_sizes.append(1280) + if len(embedding.vector_sizes) > len(hiddensizes): + embedding.tokens = [] + self.skipped_embeddings[embedding.name] = embedding + if overwrite: + shared.log.info(f"Loading Bundled embeddings: {list(data.keys())}") + for embedding in embeddings: + if embedding.name not in self.skipped_embeddings: + deref_tokenizers(embedding.tokens, tokenizers) + insert_tokens(embeddings, tokenizers) + for embedding in embeddings: + if embedding.name not in self.skipped_embeddings: + insert_vectors(embedding, tokenizers, text_encoders, hiddensizes) + self.register_embedding(embedding, shared.sd_model) + return def load_from_file(self, path, filename): name, ext = os.path.splitext(filename)