diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 878398594..0e79ae266 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -1,488 +1,488 @@ -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 lora_convert -import torch -import diffusers.models.lora -from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, sd_hijack, 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.backend == shared.Backend.DIFFUSERS: - if not hasattr(shared.sd_model, 'text_encoder') or not hasattr(shared.sd_model, 'unet'): - 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'): - 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 shared.backend != shared.Backend.DIFFUSERS: - return None - shared.sd_model.load_lora_weights(network_on_disk.filename) - 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) - if sd_module is None: - keys_failed_to_match[key_network] = key - continue - if key not in matched_networks: - matched_networks[key] = network.NetworkWeights(network_key=key_network, sd_key=key, w={}, sd_module=sd_module) - matched_networks[key].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.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx": - 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('OpenVINO: Skipping LoRa loading') - return - else: - recompile_model = True - shared.compiled_model_state.lora_model = [] - if recompile_model: - shared.compiled_model_state.lora_compile = True - sd_models.unload_model_weights(op='model') - shared.opts.cuda_compile = False - sd_models.reload_model_weights(op='model') - shared.opts.cuda_compile = True - 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: - if debug: - shared.log.debug(f'LoRA load start: name="{name}" file="{network_on_disk.filename}"') - try: - if recompile_model: - shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else 1.0}") - if shared.backend == shared.Backend.DIFFUSERS and shared.opts.lora_force_diffusers: # OpenVINO only works with Diffusers LoRa loading. - # or getattr(network_on_disk, 'shorthash', '').lower() == 'aaebf6360f7d' # sd15-lcm - # or getattr(network_on_disk, 'shorthash', '').lower() == '3d18b05e4f56' # sdxl-lcm - # or getattr(network_on_disk, 'shorthash', '').lower() == '813ea5fb1c67' # turbo sdxl-turbo - 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) - if failed_to_load_networks: - sd_hijack.model_hijack.comments.append("Networks not found: " + ", ".join(failed_to_load_networks)) - - 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") - sd_models_compile.compile_diffusers(shared.sd_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 += 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(): - global available_networks, available_network_aliases, forbidden_network_aliases, available_network_hash_lookup - available_networks.clear() - available_network_aliases.clear() - forbidden_network_aliases.clear() - available_network_hash_lookup.clear() - forbidden_network_aliases.update({"none": 1, "Addams": 1}) - os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) - directories = [] - if os.path.exists(shared.cmd_opts.lora_dir): - directories.append(shared.cmd_opts.lora_dir) - else: - shared.log.warning('LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') - if os.path.exists(shared.cmd_opts.lyco_dir): - directories.append(shared.cmd_opts.lyco_dir) - def add_network(filename): - if os.path.isdir(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 - available_network_aliases[entry.name] = entry - 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}") - - with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - for fn in files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"]): - executor.submit(add_network, fn) - print(f'Lora/LyCORIS Networks: networks={len(available_networks)} directories={directories}') - - -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 lora_convert +import torch +import diffusers.models.lora +from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, sd_hijack, 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.backend == shared.Backend.DIFFUSERS: + if not hasattr(shared.sd_model, 'text_encoder') or not hasattr(shared.sd_model, 'unet'): + 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'): + 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 shared.backend != shared.Backend.DIFFUSERS: + return None + shared.sd_model.load_lora_weights(network_on_disk.filename) + 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) + if sd_module is None: + keys_failed_to_match[key_network] = key + continue + if key not in matched_networks: + matched_networks[key] = network.NetworkWeights(network_key=key_network, sd_key=key, w={}, sd_module=sd_module) + matched_networks[key].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.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx": + 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('OpenVINO: Skipping LoRa loading') + return + else: + recompile_model = True + shared.compiled_model_state.lora_model = [] + if recompile_model: + shared.compiled_model_state.lora_compile = True + sd_models.unload_model_weights(op='model') + shared.opts.cuda_compile = False + sd_models.reload_model_weights(op='model') + shared.opts.cuda_compile = True + 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: + if debug: + shared.log.debug(f'LoRA load start: name="{name}" file="{network_on_disk.filename}"') + try: + if recompile_model: + shared.compiled_model_state.lora_model.append(f"{name}:{te_multipliers[i] if te_multipliers else 1.0}") + if shared.backend == shared.Backend.DIFFUSERS and shared.opts.lora_force_diffusers: # OpenVINO only works with Diffusers LoRa loading. + # or getattr(network_on_disk, 'shorthash', '').lower() == 'aaebf6360f7d' # sd15-lcm + # or getattr(network_on_disk, 'shorthash', '').lower() == '3d18b05e4f56' # sdxl-lcm + # or getattr(network_on_disk, 'shorthash', '').lower() == '813ea5fb1c67' # turbo sdxl-turbo + 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) + if failed_to_load_networks: + sd_hijack.model_hijack.comments.append("Networks not found: " + ", ".join(failed_to_load_networks)) + + 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") + sd_models_compile.compile_diffusers(shared.sd_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 += 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(): + global available_networks, available_network_aliases, forbidden_network_aliases, available_network_hash_lookup + available_networks.clear() + available_network_aliases.clear() + forbidden_network_aliases.clear() + available_network_hash_lookup.clear() + forbidden_network_aliases.update({"none": 1, "Addams": 1}) + os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) + directories = [] + if os.path.exists(shared.cmd_opts.lora_dir): + directories.append(shared.cmd_opts.lora_dir) + else: + shared.log.warning('LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') + if os.path.exists(shared.cmd_opts.lyco_dir): + directories.append(shared.cmd_opts.lyco_dir) + def add_network(filename): + if os.path.isdir(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 + available_network_aliases[entry.name] = entry + 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}") + + with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: + for fn in files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"]): + executor.submit(add_network, fn) + print(f'Lora/LyCORIS Networks: networks={len(available_networks)} directories={directories}') + + +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/extensions-builtin/stable-diffusion-webui-rembg b/extensions-builtin/stable-diffusion-webui-rembg index 54b723361..56a44f552 160000 --- a/extensions-builtin/stable-diffusion-webui-rembg +++ b/extensions-builtin/stable-diffusion-webui-rembg @@ -1 +1 @@ -Subproject commit 54b723361ff473a2aa4d83dca3b77f95aac1fe41 +Subproject commit 56a44f552e6adc9cf335929c16d9dfd7acd6c5bd diff --git a/modules/extensions.py b/modules/extensions.py index eeca672fd..0398e14e0 100644 --- a/modules/extensions.py +++ b/modules/extensions.py @@ -1,158 +1,158 @@ -import os -from datetime import datetime -import git -from modules import shared, errors, files_cache -from modules.paths import extensions_dir, extensions_builtin_dir - - -extensions = [] - - -if not os.path.exists(extensions_dir): - os.makedirs(extensions_dir) - - -def active(): - if shared.opts.disable_all_extensions == "all": - return [] - elif shared.opts.disable_all_extensions == "user": - return [x for x in extensions if x.enabled and x.is_builtin] - else: - return [x for x in extensions if x.enabled] - - -class Extension: - def __init__(self, name, path, enabled=True, is_builtin=False): - self.name = name - self.git_name = '' - self.path = path - self.enabled = enabled - self.status = '' - self.can_update = False - self.is_builtin = is_builtin - self.commit_hash = '' - self.commit_date = None - self.version = '' - self.description = '' - self.branch = None - self.remote = None - self.have_info_from_repo = False - self.mtime = 0 - self.ctime = 0 - - def read_info(self, force=False): - if self.have_info_from_repo and not force: - return - self.have_info_from_repo = True - repo = None - self.mtime = datetime.fromtimestamp(os.path.getmtime(self.path)).isoformat() + 'Z' - self.ctime = datetime.fromtimestamp(os.path.getctime(self.path)).isoformat() + 'Z' - try: - if os.path.exists(os.path.join(self.path, ".git")): - repo = git.Repo(self.path) - except Exception as e: - errors.display(e, f'github info from {self.path}') - if repo is None or repo.bare: - self.remote = None - else: - try: - self.status = 'unknown' - if len(repo.remotes) == 0: - shared.log.debug(f"Extension: no remotes info repo={self.name}") - return - self.git_name = repo.remotes.origin.url.split('.git')[0].split('/')[-1] - self.description = repo.description - if self.description is None or self.description.startswith("Unnamed repository"): - self.description = "[No description]" - self.remote = next(repo.remote().urls, None) - head = repo.head.commit - self.commit_date = repo.head.commit.committed_date - try: - if repo.active_branch: - self.branch = repo.active_branch.name - except Exception: - pass - self.commit_hash = head.hexsha - self.version = f"

{self.commit_hash[:8]}

{datetime.fromtimestamp(self.commit_date).strftime('%a %b%d %Y %H:%M')}

" - except Exception as ex: - shared.log.error(f"Extension: failed reading data from git repo={self.name}: {ex}") - self.remote = None - - def list_files(self, subdir, extension): - from modules import scripts - dirpath = os.path.join(self.path, subdir) - if not os.path.isdir(dirpath): - return [] - priority = '50' - if os.path.isfile(os.path.join(dirpath, "..", ".priority")): - with open(os.path.join(dirpath, "..", ".priority"), "r", encoding="utf-8") as f: - priority = str(f.read().strip()) - if priority != '50': - shared.log.debug(f'Extension priority override: {os.path.dirname(dirpath)}:{priority}') - valid_extensions = map(str.upper, ['.py','.js','.mjs']) - extension = extension.upper() - assert extension in valid_extensions, f'list_files `extension` invalid: extension={extension}, valid_extensions={valid_extensions}' - files = files_cache.list_files(dirpath, ext_filter=[extension]) - res = [scripts.ScriptFile(self.path, filename, filename, priority) for filename in sorted(files)] - return res - - def check_updates(self): - try: - repo = git.Repo(self.path) - except Exception: - self.can_update = False - return - for fetch in repo.remote().fetch(dry_run=True): - if fetch.flags != fetch.HEAD_UPTODATE: - self.can_update = True - self.status = "new commits" - return - try: - origin = repo.rev_parse('origin') - if repo.head.commit != origin: - self.can_update = True - self.status = "behind HEAD" - return - except Exception: - self.can_update = False - self.status = "unknown (remote error)" - return - self.can_update = False - self.status = "latest" - - def git_fetch(self, commit='origin'): - repo = git.Repo(self.path) - # Fix: `error: Your local changes to the following files would be overwritten by merge`, - # because WSL2 Docker set 755 file permissions instead of 644, this results to the error. - repo.git.fetch(all=True) - repo.git.reset('origin', hard=True) - repo.git.reset(commit, hard=True) - self.have_info_from_repo = False - - -def list_extensions(): - extensions.clear() - if not os.path.isdir(extensions_dir): - return - if shared.opts.disable_all_extensions == "all" or shared.opts.disable_all_extensions == "user": - shared.log.warning(f"Option set: Disable extensions: {shared.opts.disable_all_extensions}") - extension_paths = [] - extension_names = [] - extension_folders = [extensions_builtin_dir] if shared.cmd_opts.safe else [extensions_builtin_dir, extensions_dir] - for dirname in extension_folders: - if not os.path.isdir(dirname): - return - for extension_dirname in sorted(os.listdir(dirname)): - path = os.path.join(dirname, extension_dirname) - if not os.path.isdir(path): - continue - if extension_dirname in extension_names: - shared.log.info(f'Skipping conflicting extension: {path}') - continue - extension_names.append(extension_dirname) - extension_paths.append((extension_dirname, path, dirname == extensions_builtin_dir)) - disabled_extensions = shared.opts.disabled_extensions + shared.temp_disable_extensions() - for dirname, path, is_builtin in extension_paths: - extension = Extension(name=dirname, path=path, enabled=dirname not in disabled_extensions, is_builtin=is_builtin) - extensions.append(extension) - shared.log.info(f'Disabled extensions: {[e.name for e in extensions if not e.enabled]}') +import os +from datetime import datetime +import git +from modules import shared, errors, files_cache +from modules.paths import extensions_dir, extensions_builtin_dir + + +extensions = [] + + +if not os.path.exists(extensions_dir): + os.makedirs(extensions_dir) + + +def active(): + if shared.opts.disable_all_extensions == "all": + return [] + elif shared.opts.disable_all_extensions == "user": + return [x for x in extensions if x.enabled and x.is_builtin] + else: + return [x for x in extensions if x.enabled] + + +class Extension: + def __init__(self, name, path, enabled=True, is_builtin=False): + self.name = name + self.git_name = '' + self.path = path + self.enabled = enabled + self.status = '' + self.can_update = False + self.is_builtin = is_builtin + self.commit_hash = '' + self.commit_date = None + self.version = '' + self.description = '' + self.branch = None + self.remote = None + self.have_info_from_repo = False + self.mtime = 0 + self.ctime = 0 + + def read_info(self, force=False): + if self.have_info_from_repo and not force: + return + self.have_info_from_repo = True + repo = None + self.mtime = datetime.fromtimestamp(os.path.getmtime(self.path)).isoformat() + 'Z' + self.ctime = datetime.fromtimestamp(os.path.getctime(self.path)).isoformat() + 'Z' + try: + if os.path.exists(os.path.join(self.path, ".git")): + repo = git.Repo(self.path) + except Exception as e: + errors.display(e, f'github info from {self.path}') + if repo is None or repo.bare: + self.remote = None + else: + try: + self.status = 'unknown' + if len(repo.remotes) == 0: + shared.log.debug(f"Extension: no remotes info repo={self.name}") + return + self.git_name = repo.remotes.origin.url.split('.git')[0].split('/')[-1] + self.description = repo.description + if self.description is None or self.description.startswith("Unnamed repository"): + self.description = "[No description]" + self.remote = next(repo.remote().urls, None) + head = repo.head.commit + self.commit_date = repo.head.commit.committed_date + try: + if repo.active_branch: + self.branch = repo.active_branch.name + except Exception: + pass + self.commit_hash = head.hexsha + self.version = f"

{self.commit_hash[:8]}

{datetime.fromtimestamp(self.commit_date).strftime('%a %b%d %Y %H:%M')}

" + except Exception as ex: + shared.log.error(f"Extension: failed reading data from git repo={self.name}: {ex}") + self.remote = None + + def list_files(self, subdir, extension): + from modules import scripts + dirpath = os.path.join(self.path, subdir) + if not os.path.isdir(dirpath): + return [] + priority = '50' + if os.path.isfile(os.path.join(dirpath, "..", ".priority")): + with open(os.path.join(dirpath, "..", ".priority"), "r", encoding="utf-8") as f: + priority = str(f.read().strip()) + if priority != '50': + shared.log.debug(f'Extension priority override: {os.path.dirname(dirpath)}:{priority}') + valid_extensions = map(str.upper, ['.py','.js','.mjs']) + extension = extension.upper() + assert extension in valid_extensions, f'list_files `extension` invalid: extension={extension}, valid_extensions={valid_extensions}' + files = files_cache.list_files(dirpath, ext_filter=[extension]) + res = [scripts.ScriptFile(self.path, filename, filename, priority) for filename in sorted(files)] + return res + + def check_updates(self): + try: + repo = git.Repo(self.path) + except Exception: + self.can_update = False + return + for fetch in repo.remote().fetch(dry_run=True): + if fetch.flags != fetch.HEAD_UPTODATE: + self.can_update = True + self.status = "new commits" + return + try: + origin = repo.rev_parse('origin') + if repo.head.commit != origin: + self.can_update = True + self.status = "behind HEAD" + return + except Exception: + self.can_update = False + self.status = "unknown (remote error)" + return + self.can_update = False + self.status = "latest" + + def git_fetch(self, commit='origin'): + repo = git.Repo(self.path) + # Fix: `error: Your local changes to the following files would be overwritten by merge`, + # because WSL2 Docker set 755 file permissions instead of 644, this results to the error. + repo.git.fetch(all=True) + repo.git.reset('origin', hard=True) + repo.git.reset(commit, hard=True) + self.have_info_from_repo = False + + +def list_extensions(): + extensions.clear() + if not os.path.isdir(extensions_dir): + return + if shared.opts.disable_all_extensions == "all" or shared.opts.disable_all_extensions == "user": + shared.log.warning(f"Option set: Disable extensions: {shared.opts.disable_all_extensions}") + extension_paths = [] + extension_names = [] + extension_folders = [extensions_builtin_dir] if shared.cmd_opts.safe else [extensions_builtin_dir, extensions_dir] + for dirname in extension_folders: + if not os.path.isdir(dirname): + return + for extension_dirname in sorted(os.listdir(dirname)): + path = os.path.join(dirname, extension_dirname) + if not os.path.isdir(path): + continue + if extension_dirname in extension_names: + shared.log.info(f'Skipping conflicting extension: {path}') + continue + extension_names.append(extension_dirname) + extension_paths.append((extension_dirname, path, dirname == extensions_builtin_dir)) + disabled_extensions = shared.opts.disabled_extensions + shared.temp_disable_extensions() + for dirname, path, is_builtin in extension_paths: + extension = Extension(name=dirname, path=path, enabled=dirname not in disabled_extensions, is_builtin=is_builtin) + extensions.append(extension) + shared.log.info(f'Disabled extensions: {[e.name for e in extensions if not e.enabled]}') diff --git a/modules/files_cache.py b/modules/files_cache.py index c8640012f..5cd7c8104 100644 --- a/modules/files_cache.py +++ b/modules/files_cache.py @@ -1,10 +1,11 @@ -from os import scandir -import os.path as path -from typing import Dict, List, Union, Callable, Optional, Iterator -from dataclasses import dataclass, field -from installer import print_dict -from collections import UserDict import itertools +import os.path as path +from collections import UserDict +from dataclasses import dataclass, field +from os import scandir +from typing import Callable, Dict, Iterator, List, Optional, Union + +from installer import print_dict WasDirty = bool DidDelete = bool @@ -26,7 +27,7 @@ DirectoryPathIterator = Iterator[DirectoryPath] class Directory: ... - + DirectoryList = List[Directory] DirectoryIterator = Iterator[Directory] DirectoryCollection = Dict[DirectoryPath, Directory] @@ -46,7 +47,7 @@ def real_path(directory_path:DirectoryPath) -> DirectoryPath | None: @dataclass(slots=True,frozen=True) -class Directory(Directory): +class Directory(Directory): # pylint: disable=E0102 path: DirectoryPath = field(default_factory=str) @@ -67,7 +68,7 @@ class Directory(Directory): object.__setattr__(directory, 'files', dict_object.get('files')) object.__setattr__(directory, 'directories', dict_object.get('directories')) return directory - + def clear(self) -> None: self._update(Directory.from_dict({ @@ -76,14 +77,14 @@ class Directory(Directory): 'files': [], 'directories': [] })) - + def update(self, source_directory: Directory) -> Directory: if source_directory is not self: self._update(source_directory) return self - - + + def _update(self, source:Directory) -> None: assert not source.path or source.path == self.path, f'When updating a directory, the paths must match. Attemped to update Directory `{self.path}` with `{source.path}`' for dead_path in self.directories: @@ -92,7 +93,7 @@ class Directory(Directory): self.directories[:] = source.directories self.files[:] = source.files object.__setattr__(self, 'mtime', source.mtime) - + def __str__(self) -> str: return str(print_dict(self, path=self.path, mtime=self.mtime, files=len(self.files), directories=len(self.directories))) @@ -101,7 +102,7 @@ class Directory(Directory): @property def exists(self) -> DirectoryExists: return self.path and path.exists(self.path) - + @property def is_directory(self) -> IsDirectory: @@ -111,7 +112,7 @@ class Directory(Directory): @property def live_mtime(self) -> MTime: return path.getmtime(self.path) if self.is_directory else 0 - + @property def is_stale(self) -> CachedDirectoryIsStale: @@ -161,7 +162,7 @@ def get_directory(directory_or_path: DirectoryPath, /, fetch:bool=True) -> Direc return directory_or_path else: directory_or_path = directory_or_path.path - global cache_folders + global cache_folders # pylint: disable=W0602 directory_or_path = real_path(directory_or_path) if not cache_folders.get(directory_or_path, None): if fetch: @@ -247,14 +248,14 @@ def walk(top, onerror:Callable=None, /, recurse:RecursiveType=True, cached=True) def delete_cached_directory(directory_path:DirectoryPath) -> DidDelete: - global cache_folders + global cache_folders # pylint: disable=W0602 if directory_path in cache_folders: del cache_folders[directory_path] def is_directory(dir_path:DirectoryPath) -> IsDirectory: return dir_path and path.exists(dir_path) and path.isdir(dir_path) - + def directory_mtime(directory_path:DirectoryPath, /, recursive:RecursiveType=True) -> MTime: return float(max(0, *[directory.mtime for directory in get_directories(directory_path, recursive=recursive)])) @@ -263,7 +264,7 @@ def directory_mtime(directory_path:DirectoryPath, /, recursive:RecursiveType=Tru def unique_directories(directories:DirectoryPathList, /, recursive:RecursiveType=True) -> DirectoryPathIterator: '''Ensure no empty, or duplicates''' '''If we are going recursive, then directories that are children of other directories are redundant''' - directories = list(sorted(unique_paths(directories), reverse=True)) + directories = sorted(unique_paths(directories), reverse=True) #shared.log.debug(f'Directories: {directories}') while directories: directory = directories.pop() @@ -272,6 +273,7 @@ def unique_directories(directories:DirectoryPathList, /, recursive:RecursiveType if not recursive: continue _directory = path.join(directory, '') + child_directory = None while directories and directories[-1].startswith(_directory): if not callable(recursive) or not child_directory: #shared.log.debug(f'removing `{directories[-1]}` ... {_directory}') @@ -287,12 +289,9 @@ def unique_directories(directories:DirectoryPathList, /, recursive:RecursiveType else: for sub_directory in child_directory.split(path.sep): next_directory = path.join(next_directory, sub_directory) - try: - if recursive(next_directory): - _remove_directory = path.join(next_directory, '') - break - except Exception: - raise # I had thougths about suppressing the excepton, but it's probably better to not. + if recursive(next_directory): + _remove_directory = path.join(next_directory, '') + break while _remove_directory and directories: _d = directories.pop() #shared.log.info(f'Doing the while thing: {_remove_directory} - {_d}') @@ -302,14 +301,14 @@ def unique_directories(directories:DirectoryPathList, /, recursive:RecursiveType def unique_paths(directory_paths:DirectoryPathList) -> DirectoryPathIterator: return ( - key - for key - in { - real_directory_path: True - for real_directory_path + key + for key + in { + real_directory_path: True + for real_directory_path in filter(bool, [ - real_path(directory_path) - for directory_path + real_path(directory_path) + for directory_path in filter(bool, directory_paths) ]) } @@ -320,7 +319,7 @@ def get_directories(*directory_paths: DirectoryPathList, fetch:bool=True, recurs return filter( bool, ( - get_directory(directory_path, fetch=fetch) + get_directory(directory_path, fetch=fetch) for directory_path in unique_directories( directory_paths, recursive=recursive ) @@ -331,22 +330,22 @@ def get_directories(*directory_paths: DirectoryPathList, fetch:bool=True, recurs def directory_files(*directories_or_paths: DirectoryPathList|DirectoryList, recursive: RecursiveType=True) -> FilePathIterator: return itertools.chain.from_iterable( itertools.chain( - directory_object.files, + directory_object.files, [] if not recursive else itertools.chain.from_iterable( directory_files(directory, recursive=recursive) for directory in filter( - bool, + bool, map( - get_directory, + get_directory, filter( ( ( bool if recursive else False ) - if not callable(recursive) + if not callable(recursive) else recursive - ), + ), directory_object.directories ) ) @@ -386,11 +385,11 @@ def filter_files(file_paths: FilePathList, ext_filter: Optional[ExtensionList]=N def list_files(*directory_paths:DirectoryPathList, ext_filter: Optional[ExtensionList]=None, ext_blacklist: Optional[ExtensionList]=None, recursive:RecursiveType=True) -> FilePathIterator: return filter_files(itertools.chain.from_iterable( directory_files(directory, recursive=recursive) - for directory + for directory in get_directories( *directory_paths, recursive=recursive ) ), ext_filter, ext_blacklist) -cache_folders = DirectoryCache({}) \ No newline at end of file +cache_folders = DirectoryCache({}) diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index 3ca2f9207..31ed6d772 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -1,753 +1,753 @@ -import datetime -import html -import os -from collections import deque -import inspect -from statistics import stdev, mean -from rich import progress -import tqdm -import torch -from torch import einsum -from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ -from einops import rearrange, repeat -from ldm.util import default -from modules import devices, processing, sd_models, shared, hashes, sd_hijack_checkpoint, errors, files_cache -import modules.textual_inversion.dataset -from modules.textual_inversion import textual_inversion, ti_logging -from modules.textual_inversion.learn_schedule import LearnRateScheduler - - -optimizer_dict = {optim_name : cls_obj for optim_name, cls_obj in inspect.getmembers(torch.optim, inspect.isclass) if optim_name != "Optimizer"} - -class HypernetworkModule(torch.nn.Module): - activation_dict = { - "linear": torch.nn.Identity, - "relu": torch.nn.ReLU, - "leakyrelu": torch.nn.LeakyReLU, - "elu": torch.nn.ELU, - "swish": torch.nn.Hardswish, - "tanh": torch.nn.Tanh, - "sigmoid": torch.nn.Sigmoid, - } - activation_dict.update({cls_name.lower(): cls_obj for cls_name, cls_obj in inspect.getmembers(torch.nn.modules.activation) if inspect.isclass(cls_obj) and cls_obj.__module__ == 'torch.nn.modules.activation'}) - - def __init__(self, dim, state_dict=None, layer_structure=None, activation_func=None, weight_init='Normal', - add_layer_norm=False, activate_output=False, dropout_structure=None): - super().__init__() - self.multiplier = 1.0 - assert layer_structure is not None, "layer_structure must not be None" - assert layer_structure[0] == 1, "Multiplier Sequence should start with size 1!" - assert layer_structure[-1] == 1, "Multiplier Sequence should end with size 1!" - linears = [] - for i in range(len(layer_structure) - 1): - # Add a fully-connected layer - linears.append(torch.nn.Linear(int(dim * layer_structure[i]), int(dim * layer_structure[i+1]))) - # Add an activation func except last layer - if activation_func == "linear" or activation_func is None or (i >= len(layer_structure) - 2 and not activate_output): - pass - elif activation_func in self.activation_dict: - linears.append(self.activation_dict[activation_func]()) - else: - raise RuntimeError(f'hypernetwork uses an unsupported activation function: {activation_func}') - # Add layer normalization - if add_layer_norm: - linears.append(torch.nn.LayerNorm(int(dim * layer_structure[i+1]))) - # Everything should be now parsed into dropout structure, and applied here. - # Since we only have dropouts after layers, dropout structure should start with 0 and end with 0. - if dropout_structure is not None and dropout_structure[i+1] > 0: - assert 0 < dropout_structure[i+1] < 1, "Dropout probability should be 0 or float between 0 and 1!" - linears.append(torch.nn.Dropout(p=dropout_structure[i+1])) - # Code explanation : [1, 2, 1] -> dropout is missing when last_layer_dropout is false. [1, 2, 2, 1] -> [0, 0.3, 0, 0], when its True, [0, 0.3, 0.3, 0]. - self.linear = torch.nn.Sequential(*linears) - if state_dict is not None: - self.fix_old_state_dict(state_dict) - self.load_state_dict(state_dict) - else: - for layer in self.linear: - if type(layer) == torch.nn.Linear or type(layer) == torch.nn.LayerNorm: - w, b = layer.weight.data, layer.bias.data - if weight_init == "Normal" or type(layer) == torch.nn.LayerNorm: - normal_(w, mean=0.0, std=0.01) - normal_(b, mean=0.0, std=0) - elif weight_init == 'XavierUniform': - xavier_uniform_(w) - zeros_(b) - elif weight_init == 'XavierNormal': - xavier_normal_(w) - zeros_(b) - elif weight_init == 'KaimingUniform': - kaiming_uniform_(w, nonlinearity='leaky_relu' if 'leakyrelu' == activation_func else 'relu') - zeros_(b) - elif weight_init == 'KaimingNormal': - kaiming_normal_(w, nonlinearity='leaky_relu' if 'leakyrelu' == activation_func else 'relu') - zeros_(b) - else: - raise KeyError(f"Key {weight_init} is not defined as initialization!") - self.to(devices.device) - - def fix_old_state_dict(self, state_dict): - changes = { - 'linear1.bias': 'linear.0.bias', - 'linear1.weight': 'linear.0.weight', - 'linear2.bias': 'linear.1.bias', - 'linear2.weight': 'linear.1.weight', - } - for fr, to in changes.items(): - x = state_dict.get(fr, None) - if x is None: - continue - del state_dict[fr] - state_dict[to] = x - - def forward(self, x): - return x + self.linear(x) * (self.multiplier if not self.training else 1) - - def trainables(self): - layer_structure = [] - for layer in self.linear: - if type(layer) == torch.nn.Linear or type(layer) == torch.nn.LayerNorm: - layer_structure += [layer.weight, layer.bias] - return layer_structure - - -#param layer_structure : sequence used for length, use_dropout : controlling boolean, last_layer_dropout : for compatibility check. -def parse_dropout_structure(layer_structure, use_dropout, last_layer_dropout): - if layer_structure is None: - layer_structure = [1, 2, 1] - if not use_dropout: - return [0] * len(layer_structure) - dropout_values = [0] - dropout_values.extend([0.3] * (len(layer_structure) - 3)) - if last_layer_dropout: - dropout_values.append(0.3) - else: - dropout_values.append(0) - dropout_values.append(0) - return dropout_values - - -class Hypernetwork: - filename = None - name = None - - def __init__(self, name=None, enable_sizes=None, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, activate_output=False, **kwargs): - self.filename = None - self.name = name - self.layers = {} - self.step = 0 - self.sd_checkpoint = None - self.sd_checkpoint_name = None - self.layer_structure = layer_structure - self.activation_func = activation_func - self.weight_init = weight_init - self.add_layer_norm = add_layer_norm - self.use_dropout = use_dropout - self.activate_output = activate_output - self.last_layer_dropout = kwargs.get('last_layer_dropout', True) - self.dropout_structure = kwargs.get('dropout_structure', None) - if self.dropout_structure is None: - self.dropout_structure = parse_dropout_structure(self.layer_structure, self.use_dropout, self.last_layer_dropout) - self.optimizer_name = None - self.optimizer_state_dict = None - self.optional_info = None - for size in enable_sizes or []: - self.layers[size] = ( - HypernetworkModule(size, None, self.layer_structure, self.activation_func, self.weight_init, - self.add_layer_norm, self.activate_output, dropout_structure=self.dropout_structure), - HypernetworkModule(size, None, self.layer_structure, self.activation_func, self.weight_init, - self.add_layer_norm, self.activate_output, dropout_structure=self.dropout_structure), - ) - self.eval() - - def weights(self): - res = [] - for layers in self.layers.values(): - for layer in layers: - res += layer.parameters() - return res - - def train(self, mode=True): - for layers in self.layers.values(): - for layer in layers: - layer.train(mode=mode) - for param in layer.parameters(): - param.requires_grad = mode - - def to(self, device): - for layers in self.layers.values(): - for layer in layers: - layer.to(device) - - return self - - def set_multiplier(self, multiplier): - for layers in self.layers.values(): - for layer in layers: - layer.multiplier = multiplier - - return self - - def eval(self): - for layers in self.layers.values(): - for layer in layers: - layer.eval() - for param in layer.parameters(): - param.requires_grad = False - - def save(self, filename): - state_dict = {} - optimizer_saved_dict = {} - for k, v in self.layers.items(): - state_dict[k] = (v[0].state_dict(), v[1].state_dict()) - state_dict['step'] = self.step - state_dict['name'] = self.name - state_dict['layer_structure'] = self.layer_structure - state_dict['activation_func'] = self.activation_func - state_dict['is_layer_norm'] = self.add_layer_norm - state_dict['weight_initialization'] = self.weight_init - state_dict['sd_checkpoint'] = self.sd_checkpoint - state_dict['sd_checkpoint_name'] = self.sd_checkpoint_name - state_dict['activate_output'] = self.activate_output - state_dict['use_dropout'] = self.use_dropout - state_dict['dropout_structure'] = self.dropout_structure - state_dict['last_layer_dropout'] = (self.dropout_structure[-2] != 0) if self.dropout_structure is not None else self.last_layer_dropout - state_dict['optional_info'] = self.optional_info if self.optional_info else None - if self.optimizer_name is not None: - optimizer_saved_dict['optimizer_name'] = self.optimizer_name - torch.save(state_dict, filename) - if shared.opts.save_optimizer_state and self.optimizer_state_dict: - optimizer_saved_dict['hash'] = self.shorthash() - optimizer_saved_dict['optimizer_state_dict'] = self.optimizer_state_dict - torch.save(optimizer_saved_dict, f"{filename}.optim") - - def load(self, filename): - self.filename = filename if os.path.exists(filename) else os.path.join(shared.opts.hypernetwork_dir, filename) - if self.name is None: - self.name = os.path.splitext(os.path.basename(self.filename))[0] - with progress.open(self.filename, 'rb', description=f'Load hypernetwork: [cyan]{self.filename}', auto_refresh=True, console=shared.console) as f: - state_dict = torch.load(f, map_location='cpu') - self.layer_structure = state_dict.get('layer_structure', [1, 2, 1]) - self.optional_info = state_dict.get('optional_info', None) - self.activation_func = state_dict.get('activation_func', None) - self.weight_init = state_dict.get('weight_initialization', 'Normal') - self.add_layer_norm = state_dict.get('is_layer_norm', False) - self.dropout_structure = state_dict.get('dropout_structure', None) - self.use_dropout = True if self.dropout_structure is not None and any(self.dropout_structure) else state_dict.get('use_dropout', False) - self.activate_output = state_dict.get('activate_output', True) - self.last_layer_dropout = state_dict.get('last_layer_dropout', False) - # Dropout structure should have same length as layer structure, Every digits should be in [0,1), and last digit must be 0. - if self.dropout_structure is None: - self.dropout_structure = parse_dropout_structure(self.layer_structure, self.use_dropout, self.last_layer_dropout) - if shared.opts.print_hypernet_extra: - if self.optional_info is not None: - print(f" INFO:\n {self.optional_info}\n") - print(f" Layer structure: {self.layer_structure}") - print(f" Activation function: {self.activation_func}") - print(f" Weight initialization: {self.weight_init}") - print(f" Layer norm: {self.add_layer_norm}") - print(f" Dropout usage: {self.use_dropout}" ) - print(f" Activate last layer: {self.activate_output}") - print(f" Dropout structure: {self.dropout_structure}") - optimizer_saved_dict = torch.load(self.filename + '.optim', map_location='cpu') if os.path.exists(self.filename + '.optim') else {} - if self.shorthash() == optimizer_saved_dict.get('hash', None): - self.optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None) - else: - self.optimizer_state_dict = None - if self.optimizer_state_dict: - self.optimizer_name = optimizer_saved_dict.get('optimizer_name', 'AdamW') - if shared.opts.print_hypernet_extra: - print("Load existing optimizer from checkpoint") - print(f"Optimizer name is {self.optimizer_name}") - else: - self.optimizer_name = "AdamW" - if shared.opts.print_hypernet_extra: - print("No saved optimizer exists in checkpoint") - for size, sd in state_dict.items(): - if type(size) == int: - self.layers[size] = ( - HypernetworkModule(size, sd[0], self.layer_structure, self.activation_func, self.weight_init, - self.add_layer_norm, self.activate_output, self.dropout_structure), - HypernetworkModule(size, sd[1], self.layer_structure, self.activation_func, self.weight_init, - self.add_layer_norm, self.activate_output, self.dropout_structure), - ) - self.name = state_dict.get('name', self.name) - self.step = state_dict.get('step', 0) - self.sd_checkpoint = state_dict.get('sd_checkpoint', None) - self.sd_checkpoint_name = state_dict.get('sd_checkpoint_name', None) - self.eval() - - def shorthash(self): - sha256 = hashes.sha256(self.filename, f'hypernet/{self.name}') - return sha256[0:10] if sha256 else None - - -def list_hypernetworks(path): - hypernetworks = { - os.path.splitext(os.path.basename(hypernetwork_path))[0]: hypernetwork_path - for hypernetwork_path - in files_cache.list_files( - path, - ext_filter=['.pt'], - recursive=files_cache.not_hidden - ) - } - return hypernetworks - - -def load_hypernetwork(name): - path = shared.hypernetworks.get(name, None) - if path is None: - return None - hypernetwork = Hypernetwork() - try: - hypernetwork.load(path) - except Exception as e: - errors.display(e, f'hypernetwork load: {path}') - return None - return hypernetwork - - -def load_hypernetworks(names, multipliers=None): - already_loaded = {} - for hypernetwork in shared.loaded_hypernetworks: - if hypernetwork.name in names: - already_loaded[hypernetwork.name] = hypernetwork - shared.loaded_hypernetworks.clear() - for i, name in enumerate(names): - hypernetwork = already_loaded.get(name, None) - if hypernetwork is None: - hypernetwork = load_hypernetwork(name) - if hypernetwork is None: - continue - hypernetwork.set_multiplier(multipliers[i] if multipliers else 1.0) - shared.loaded_hypernetworks.append(hypernetwork) - - -def find_closest_hypernetwork_name(search: str): - if not search: - return None - search = search.lower() - applicable = [name for name in shared.hypernetworks if search in name.lower()] - if not applicable: - return None - applicable = sorted(applicable, key=lambda name: len(name)) - return applicable[0] - - -def apply_single_hypernetwork(hypernetwork, context_k, context_v, layer=None): - hypernetwork_layers = (hypernetwork.layers if hypernetwork is not None else {}).get(context_k.shape[2], None) - if hypernetwork_layers is None: - return context_k, context_v - if layer is not None: - layer.hyper_k = hypernetwork_layers[0] - layer.hyper_v = hypernetwork_layers[1] - context_k = devices.cond_cast_unet(hypernetwork_layers[0](devices.cond_cast_float(context_k))) - context_v = devices.cond_cast_unet(hypernetwork_layers[1](devices.cond_cast_float(context_v))) - return context_k, context_v - - -def apply_hypernetworks(hypernetworks, context, layer=None): - context_k = context - context_v = context - for hypernetwork in hypernetworks: - context_k, context_v = apply_single_hypernetwork(hypernetwork, context_k, context_v, layer) - return context_k, context_v - - -def attention_CrossAttention_forward(self, x, context=None, mask=None): - h = self.heads - q = self.to_q(x) - context = default(context, x) - context_k, context_v = apply_hypernetworks(shared.loaded_hypernetworks, context, self) - k = self.to_k(context_k) - v = self.to_v(context_v) - q, k, v = (rearrange(t, 'b n (h d) -> (b h) n d', h=h) for t in (q, k, v)) - sim = einsum('b i d, b j d -> b i j', q, k) * self.scale - if mask is not None: - mask = rearrange(mask, 'b ... -> b (...)') - max_neg_value = -torch.finfo(sim.dtype).max - mask = repeat(mask, 'b j -> (b h) () j', h=h) - sim.masked_fill_(~mask, max_neg_value) - # attention, what we cannot get enough of - attn = sim.softmax(dim=-1) - out = einsum('b i j, b j d -> b i d', attn, v) - out = rearrange(out, '(b h) n d -> b n (h d)', h=h) - return self.to_out(out) - - -def stack_conds(conds): - if len(conds) == 1: - return torch.stack(conds) - # same as in reconstruct_multicond_batch - token_count = max([x.shape[0] for x in conds]) - for i in range(len(conds)): - if conds[i].shape[0] != token_count: - last_vector = conds[i][-1:] - last_vector_repeated = last_vector.repeat([token_count - conds[i].shape[0], 1]) - conds[i] = torch.vstack([conds[i], last_vector_repeated]) - return torch.stack(conds) - - -def statistics(data): - if len(data) < 2: - std = 0 - else: - std = stdev(data) - total_information = f"loss:{mean(data):.3f}" + "\u00B1" + f"({std/ (len(data) ** 0.5):.3f})" - recent_data = data[-32:] - if len(recent_data) < 2: - std = 0 - else: - std = stdev(recent_data) - recent_information = f"recent 32 loss:{mean(recent_data):.3f}" + "\u00B1" + f"({std / (len(recent_data) ** 0.5):.3f})" - return total_information, recent_information - - -def report_statistics(loss_info:dict): - keys = sorted(loss_info.keys(), key=lambda x: sum(loss_info[x]) / len(loss_info[x])) - for key in keys: - try: - print("Loss statistics for file " + key) - info, recent = statistics(list(loss_info[key])) - print(info) - print(recent) - except Exception as e: - print(e) - - -def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, dropout_structure=None): - # Remove illegal characters from name. - name = "".join( x for x in name if (x.isalnum() or x in "._- ")) - assert name, "Name cannot be empty!" - fn = os.path.join(shared.opts.hypernetwork_dir, f"{name}.pt") - if not overwrite_old: - assert not os.path.exists(fn), f"file {fn} already exists" - if type(layer_structure) == str: - layer_structure = [float(x.strip()) for x in layer_structure.split(",")] - if use_dropout and dropout_structure and type(dropout_structure) == str: - dropout_structure = [float(x.strip()) for x in dropout_structure.split(",")] - else: - dropout_structure = [0] * len(layer_structure) - hypernet = modules.hypernetworks.hypernetwork.Hypernetwork( - name=name, - enable_sizes=[int(x) for x in enable_sizes], - layer_structure=layer_structure, - activation_func=activation_func, - weight_init=weight_init, - add_layer_norm=add_layer_norm, - use_dropout=use_dropout, - dropout_structure=dropout_structure - ) - hypernet.save(fn) - shared.reload_hypernetworks() - return name - - -def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_hypernetwork_every, template_filename, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument - # images allows training previews to have infotext. Importing it at the top causes a circular import problem. - from modules import images - - save_hypernetwork_every = save_hypernetwork_every or 0 - create_image_every = create_image_every or 0 - template_file = textual_inversion.textual_inversion_templates.get(template_filename, None) - textual_inversion.validate_train_inputs(hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_hypernetwork_every, create_image_every, name="hypernetwork") - template_file = template_file.path - - path = shared.hypernetworks.get(hypernetwork_name, None) - hypernetwork = Hypernetwork() - hypernetwork.load(path) - shared.loaded_hypernetworks = [hypernetwork] - - shared.state.job = "train" - shared.state.textinfo = "Initializing hypernetwork training..." - shared.state.job_count = steps - - hypernetwork_name = hypernetwork_name.rsplit('(', 1)[0] - filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt') - - log_directory = os.path.join(log_directory, datetime.datetime.now().strftime("%Y-%m-%d"), hypernetwork_name) - unload = shared.opts.unload_models_when_training - - if save_hypernetwork_every > 0: - hypernetwork_dir = os.path.join(log_directory, "hypernetworks") - os.makedirs(hypernetwork_dir, exist_ok=True) - else: - hypernetwork_dir = None - - if create_image_every > 0: - images_dir = os.path.join(log_directory, "images") - os.makedirs(images_dir, exist_ok=True) - else: - images_dir = None - - checkpoint = sd_models.select_checkpoint() - - initial_step = hypernetwork.step or 0 - if initial_step >= steps: - shared.state.textinfo = "Model has already been trained beyond specified max steps" - return hypernetwork, filename - - scheduler = LearnRateScheduler(learn_rate, steps, initial_step) - - clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else None - if clip_grad: - clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False) - - if shared.opts.training_enable_tensorboard: - tensorboard_writer = textual_inversion.tensorboard_setup(log_directory) - - # dataset loading may take a while, so input validations and early returns should be done before this - shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..." - - pin_memory = shared.opts.pin_memory - - ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=hypernetwork_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, include_cond=True, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight) - - if shared.opts.save_training_settings_to_txt: - saved_params = dict( - model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds), - **{field: getattr(hypernetwork, field) for field in ['layer_structure', 'activation_func', 'weight_init', 'add_layer_norm', 'use_dropout', ]} - ) - ti_logging.save_settings_to_file(log_directory, {**saved_params, **locals()}) - - latent_sampling_method = ds.latent_sampling_method - - dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory) - - old_parallel_processing_allowed = shared.parallel_processing_allowed - - if unload: - shared.parallel_processing_allowed = False - shared.sd_model.cond_stage_model.to(devices.cpu) - shared.sd_model.first_stage_model.to(devices.cpu) - - weights = hypernetwork.weights() - hypernetwork.train() - - # Here we use optimizer from saved HN, or we can specify as UI option. - if hypernetwork.optimizer_name in optimizer_dict: - optimizer = optimizer_dict[hypernetwork.optimizer_name](params=weights, lr=scheduler.learn_rate) - optimizer_name = hypernetwork.optimizer_name - else: - print(f"Optimizer type {hypernetwork.optimizer_name} is not defined!") - optimizer = torch.optim.AdamW(params=weights, lr=scheduler.learn_rate) - optimizer_name = 'AdamW' - - if hypernetwork.optimizer_state_dict: # This line must be changed if Optimizer type can be different from saved optimizer. - try: - optimizer.load_state_dict(hypernetwork.optimizer_state_dict) - except RuntimeError as e: - print("Cannot resume from saved optimizer!") - print(e) - - scaler = torch.cuda.amp.GradScaler() - - batch_size = ds.batch_size - gradient_step = ds.gradient_step - # n steps = batch_size * gradient_step * n image processed - steps_per_epoch = len(ds) // batch_size // gradient_step - max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step - loss_step = 0 - _loss_step = 0 #internal - # size = len(ds.indexes) - # loss_dict = defaultdict(lambda : deque(maxlen = 1024)) - loss_logging = deque(maxlen=len(ds) * 3) # this should be configurable parameter, this is 3 * epoch(dataset size) - # losses = torch.zeros((size,)) - # previous_mean_losses = [0] - # previous_mean_loss = 0 - # print("Mean loss of {} elements".format(size)) - - _steps_without_grad = 0 - - last_saved_file = "" - last_saved_image = "" - forced_filename = "" - - pbar = tqdm.tqdm(total=steps - initial_step) - try: - sd_hijack_checkpoint.add() - - for _i in range((steps-initial_step) * gradient_step): - if scheduler.finished: - break - if shared.state.interrupted: - break - for j, batch in enumerate(dl): - # works as a drop_last=True for gradient accumulation - if j == max_steps_per_epoch: - break - scheduler.apply(optimizer, hypernetwork.step) - if scheduler.finished: - break - if shared.state.interrupted: - break - - if clip_grad: - clip_grad_sched.step(hypernetwork.step) - - with devices.autocast(): - x = batch.latent_sample.to(devices.device, non_blocking=pin_memory) - if use_weight: - w = batch.weight.to(devices.device, non_blocking=pin_memory) - if tag_drop_out != 0 or shuffle_tags: - shared.sd_model.cond_stage_model.to(devices.device) - c = shared.sd_model.cond_stage_model(batch.cond_text).to(devices.device, non_blocking=pin_memory) - shared.sd_model.cond_stage_model.to(devices.cpu) - else: - c = stack_conds(batch.cond).to(devices.device, non_blocking=pin_memory) - if use_weight: - loss = shared.sd_model.weighted_forward(x, c, w)[0] / gradient_step - del w - else: - loss = shared.sd_model.forward(x, c)[0] / gradient_step - del x - del c - _loss_step += loss.item() - - scaler.scale(loss).backward() - # go back until we reach gradient accumulation steps - if (j + 1) % gradient_step != 0: - continue - loss_logging.append(_loss_step) - if clip_grad: - clip_grad(weights, clip_grad_sched.learn_rate) - - scaler.step(optimizer) - scaler.update() - hypernetwork.step += 1 - pbar.update() - optimizer.zero_grad(set_to_none=True) - loss_step = _loss_step - _loss_step = 0 - steps_done = hypernetwork.step + 1 - epoch_num = hypernetwork.step // steps_per_epoch - epoch_step = hypernetwork.step % steps_per_epoch - - description = f"Training hypernetwork [Epoch {epoch_num}: {epoch_step+1}/{steps_per_epoch}]loss: {loss_step:.7f}" - pbar.set_description(description) - if hypernetwork_dir is not None and steps_done % save_hypernetwork_every == 0: - # Before saving, change name to match current checkpoint. - hypernetwork_name_every = f'{hypernetwork_name}-{steps_done}' - last_saved_file = os.path.join(hypernetwork_dir, f'{hypernetwork_name_every}.pt') - hypernetwork.optimizer_name = optimizer_name - if shared.opts.save_optimizer_state: - hypernetwork.optimizer_state_dict = optimizer.state_dict() - save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, last_saved_file) - hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory. - - - - if shared.opts.training_enable_tensorboard: - epoch_num = hypernetwork.step // len(ds) - epoch_step = hypernetwork.step - (epoch_num * len(ds)) + 1 - mean_loss = sum(loss_logging) / len(loss_logging) - textual_inversion.tensorboard_add(tensorboard_writer, loss=mean_loss, global_step=hypernetwork.step, step=epoch_step, learn_rate=scheduler.learn_rate, epoch_num=epoch_num) - - textual_inversion.write_loss(log_directory, "hypernetwork_loss.csv", hypernetwork.step, steps_per_epoch, { - "loss": f"{loss_step:.7f}", - "learn_rate": scheduler.learn_rate - }) - - if images_dir is not None and steps_done % create_image_every == 0: - forced_filename = f'{hypernetwork_name}-{steps_done}' - last_saved_image = os.path.join(images_dir, forced_filename) - hypernetwork.eval() - rng_state = torch.get_rng_state() - cuda_rng_state = None - cuda_rng_state = torch.cuda.get_rng_state_all() - shared.sd_model.cond_stage_model.to(devices.device) - shared.sd_model.first_stage_model.to(devices.device) - - p = processing.StableDiffusionProcessingTxt2Img( - sd_model=shared.sd_model, - do_not_save_grid=True, - do_not_save_samples=True, - ) - - p.disable_extra_networks = True - - if preview_from_txt2img: - p.prompt = preview_prompt - p.negative_prompt = preview_negative_prompt - p.steps = preview_steps - p.sampler_name = processing.get_sampler_name(preview_sampler_index) - p.cfg_scale = preview_cfg_scale - p.seed = preview_seed - p.width = preview_width - p.height = preview_height - else: - p.prompt = batch.cond_text[0] - p.steps = 20 - p.width = training_width - p.height = training_height - - preview_text = p.prompt - - processed = processing.process_images(p) - image = processed.images[0] if len(processed.images) > 0 else None - - if unload: - shared.sd_model.cond_stage_model.to(devices.cpu) - shared.sd_model.first_stage_model.to(devices.cpu) - torch.set_rng_state(rng_state) - torch.cuda.set_rng_state_all(cuda_rng_state) - hypernetwork.train() - if image is not None: - shared.state.assign_current_image(image) - if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images: - textual_inversion.tensorboard_add_image(tensorboard_writer, - f"Validation at epoch {epoch_num}", image, - hypernetwork.step) - last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) - last_saved_image += f", prompt: {preview_text}" - - shared.state.job_no = hypernetwork.step - - shared.state.textinfo = f""" -

-Loss: {loss_step:.7f}
-Step: {steps_done}
-Last prompt: {html.escape(batch.cond_text[0])}
-Last saved hypernetwork: {html.escape(last_saved_file)}
-Last saved image: {html.escape(last_saved_image)}
-

-""" - except Exception as e: - errors.display(e, 'hypernetwork train') - finally: - pbar.leave = False - pbar.close() - hypernetwork.eval() - #report_statistics(loss_dict) - sd_hijack_checkpoint.remove() - - - - filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt') - hypernetwork.optimizer_name = optimizer_name - if shared.opts.save_optimizer_state: - hypernetwork.optimizer_state_dict = optimizer.state_dict() - save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename) - - del optimizer - hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory. - shared.sd_model.cond_stage_model.to(devices.device) - shared.sd_model.first_stage_model.to(devices.device) - shared.parallel_processing_allowed = old_parallel_processing_allowed - - return hypernetwork, filename - -def save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename): - old_hypernetwork_name = hypernetwork.name - old_sd_checkpoint = hypernetwork.sd_checkpoint if hasattr(hypernetwork, "sd_checkpoint") else None - old_sd_checkpoint_name = hypernetwork.sd_checkpoint_name if hasattr(hypernetwork, "sd_checkpoint_name") else None - try: - hypernetwork.sd_checkpoint = checkpoint.shorthash - hypernetwork.sd_checkpoint_name = checkpoint.model_name - hypernetwork.name = hypernetwork_name - hypernetwork.save(filename) - except Exception: - hypernetwork.sd_checkpoint = old_sd_checkpoint - hypernetwork.sd_checkpoint_name = old_sd_checkpoint_name - hypernetwork.name = old_hypernetwork_name - raise +import datetime +import html +import os +from collections import deque +import inspect +from statistics import stdev, mean +from rich import progress +import tqdm +import torch +from torch import einsum +from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ +from einops import rearrange, repeat +from ldm.util import default +from modules import devices, processing, sd_models, shared, hashes, sd_hijack_checkpoint, errors, files_cache +import modules.textual_inversion.dataset +from modules.textual_inversion import textual_inversion, ti_logging +from modules.textual_inversion.learn_schedule import LearnRateScheduler + + +optimizer_dict = {optim_name : cls_obj for optim_name, cls_obj in inspect.getmembers(torch.optim, inspect.isclass) if optim_name != "Optimizer"} + +class HypernetworkModule(torch.nn.Module): + activation_dict = { + "linear": torch.nn.Identity, + "relu": torch.nn.ReLU, + "leakyrelu": torch.nn.LeakyReLU, + "elu": torch.nn.ELU, + "swish": torch.nn.Hardswish, + "tanh": torch.nn.Tanh, + "sigmoid": torch.nn.Sigmoid, + } + activation_dict.update({cls_name.lower(): cls_obj for cls_name, cls_obj in inspect.getmembers(torch.nn.modules.activation) if inspect.isclass(cls_obj) and cls_obj.__module__ == 'torch.nn.modules.activation'}) + + def __init__(self, dim, state_dict=None, layer_structure=None, activation_func=None, weight_init='Normal', + add_layer_norm=False, activate_output=False, dropout_structure=None): + super().__init__() + self.multiplier = 1.0 + assert layer_structure is not None, "layer_structure must not be None" + assert layer_structure[0] == 1, "Multiplier Sequence should start with size 1!" + assert layer_structure[-1] == 1, "Multiplier Sequence should end with size 1!" + linears = [] + for i in range(len(layer_structure) - 1): + # Add a fully-connected layer + linears.append(torch.nn.Linear(int(dim * layer_structure[i]), int(dim * layer_structure[i+1]))) + # Add an activation func except last layer + if activation_func == "linear" or activation_func is None or (i >= len(layer_structure) - 2 and not activate_output): + pass + elif activation_func in self.activation_dict: + linears.append(self.activation_dict[activation_func]()) + else: + raise RuntimeError(f'hypernetwork uses an unsupported activation function: {activation_func}') + # Add layer normalization + if add_layer_norm: + linears.append(torch.nn.LayerNorm(int(dim * layer_structure[i+1]))) + # Everything should be now parsed into dropout structure, and applied here. + # Since we only have dropouts after layers, dropout structure should start with 0 and end with 0. + if dropout_structure is not None and dropout_structure[i+1] > 0: + assert 0 < dropout_structure[i+1] < 1, "Dropout probability should be 0 or float between 0 and 1!" + linears.append(torch.nn.Dropout(p=dropout_structure[i+1])) + # Code explanation : [1, 2, 1] -> dropout is missing when last_layer_dropout is false. [1, 2, 2, 1] -> [0, 0.3, 0, 0], when its True, [0, 0.3, 0.3, 0]. + self.linear = torch.nn.Sequential(*linears) + if state_dict is not None: + self.fix_old_state_dict(state_dict) + self.load_state_dict(state_dict) + else: + for layer in self.linear: + if type(layer) == torch.nn.Linear or type(layer) == torch.nn.LayerNorm: + w, b = layer.weight.data, layer.bias.data + if weight_init == "Normal" or type(layer) == torch.nn.LayerNorm: + normal_(w, mean=0.0, std=0.01) + normal_(b, mean=0.0, std=0) + elif weight_init == 'XavierUniform': + xavier_uniform_(w) + zeros_(b) + elif weight_init == 'XavierNormal': + xavier_normal_(w) + zeros_(b) + elif weight_init == 'KaimingUniform': + kaiming_uniform_(w, nonlinearity='leaky_relu' if 'leakyrelu' == activation_func else 'relu') + zeros_(b) + elif weight_init == 'KaimingNormal': + kaiming_normal_(w, nonlinearity='leaky_relu' if 'leakyrelu' == activation_func else 'relu') + zeros_(b) + else: + raise KeyError(f"Key {weight_init} is not defined as initialization!") + self.to(devices.device) + + def fix_old_state_dict(self, state_dict): + changes = { + 'linear1.bias': 'linear.0.bias', + 'linear1.weight': 'linear.0.weight', + 'linear2.bias': 'linear.1.bias', + 'linear2.weight': 'linear.1.weight', + } + for fr, to in changes.items(): + x = state_dict.get(fr, None) + if x is None: + continue + del state_dict[fr] + state_dict[to] = x + + def forward(self, x): + return x + self.linear(x) * (self.multiplier if not self.training else 1) + + def trainables(self): + layer_structure = [] + for layer in self.linear: + if type(layer) == torch.nn.Linear or type(layer) == torch.nn.LayerNorm: + layer_structure += [layer.weight, layer.bias] + return layer_structure + + +#param layer_structure : sequence used for length, use_dropout : controlling boolean, last_layer_dropout : for compatibility check. +def parse_dropout_structure(layer_structure, use_dropout, last_layer_dropout): + if layer_structure is None: + layer_structure = [1, 2, 1] + if not use_dropout: + return [0] * len(layer_structure) + dropout_values = [0] + dropout_values.extend([0.3] * (len(layer_structure) - 3)) + if last_layer_dropout: + dropout_values.append(0.3) + else: + dropout_values.append(0) + dropout_values.append(0) + return dropout_values + + +class Hypernetwork: + filename = None + name = None + + def __init__(self, name=None, enable_sizes=None, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, activate_output=False, **kwargs): + self.filename = None + self.name = name + self.layers = {} + self.step = 0 + self.sd_checkpoint = None + self.sd_checkpoint_name = None + self.layer_structure = layer_structure + self.activation_func = activation_func + self.weight_init = weight_init + self.add_layer_norm = add_layer_norm + self.use_dropout = use_dropout + self.activate_output = activate_output + self.last_layer_dropout = kwargs.get('last_layer_dropout', True) + self.dropout_structure = kwargs.get('dropout_structure', None) + if self.dropout_structure is None: + self.dropout_structure = parse_dropout_structure(self.layer_structure, self.use_dropout, self.last_layer_dropout) + self.optimizer_name = None + self.optimizer_state_dict = None + self.optional_info = None + for size in enable_sizes or []: + self.layers[size] = ( + HypernetworkModule(size, None, self.layer_structure, self.activation_func, self.weight_init, + self.add_layer_norm, self.activate_output, dropout_structure=self.dropout_structure), + HypernetworkModule(size, None, self.layer_structure, self.activation_func, self.weight_init, + self.add_layer_norm, self.activate_output, dropout_structure=self.dropout_structure), + ) + self.eval() + + def weights(self): + res = [] + for layers in self.layers.values(): + for layer in layers: + res += layer.parameters() + return res + + def train(self, mode=True): + for layers in self.layers.values(): + for layer in layers: + layer.train(mode=mode) + for param in layer.parameters(): + param.requires_grad = mode + + def to(self, device): + for layers in self.layers.values(): + for layer in layers: + layer.to(device) + + return self + + def set_multiplier(self, multiplier): + for layers in self.layers.values(): + for layer in layers: + layer.multiplier = multiplier + + return self + + def eval(self): + for layers in self.layers.values(): + for layer in layers: + layer.eval() + for param in layer.parameters(): + param.requires_grad = False + + def save(self, filename): + state_dict = {} + optimizer_saved_dict = {} + for k, v in self.layers.items(): + state_dict[k] = (v[0].state_dict(), v[1].state_dict()) + state_dict['step'] = self.step + state_dict['name'] = self.name + state_dict['layer_structure'] = self.layer_structure + state_dict['activation_func'] = self.activation_func + state_dict['is_layer_norm'] = self.add_layer_norm + state_dict['weight_initialization'] = self.weight_init + state_dict['sd_checkpoint'] = self.sd_checkpoint + state_dict['sd_checkpoint_name'] = self.sd_checkpoint_name + state_dict['activate_output'] = self.activate_output + state_dict['use_dropout'] = self.use_dropout + state_dict['dropout_structure'] = self.dropout_structure + state_dict['last_layer_dropout'] = (self.dropout_structure[-2] != 0) if self.dropout_structure is not None else self.last_layer_dropout + state_dict['optional_info'] = self.optional_info if self.optional_info else None + if self.optimizer_name is not None: + optimizer_saved_dict['optimizer_name'] = self.optimizer_name + torch.save(state_dict, filename) + if shared.opts.save_optimizer_state and self.optimizer_state_dict: + optimizer_saved_dict['hash'] = self.shorthash() + optimizer_saved_dict['optimizer_state_dict'] = self.optimizer_state_dict + torch.save(optimizer_saved_dict, f"{filename}.optim") + + def load(self, filename): + self.filename = filename if os.path.exists(filename) else os.path.join(shared.opts.hypernetwork_dir, filename) + if self.name is None: + self.name = os.path.splitext(os.path.basename(self.filename))[0] + with progress.open(self.filename, 'rb', description=f'Load hypernetwork: [cyan]{self.filename}', auto_refresh=True, console=shared.console) as f: + state_dict = torch.load(f, map_location='cpu') + self.layer_structure = state_dict.get('layer_structure', [1, 2, 1]) + self.optional_info = state_dict.get('optional_info', None) + self.activation_func = state_dict.get('activation_func', None) + self.weight_init = state_dict.get('weight_initialization', 'Normal') + self.add_layer_norm = state_dict.get('is_layer_norm', False) + self.dropout_structure = state_dict.get('dropout_structure', None) + self.use_dropout = True if self.dropout_structure is not None and any(self.dropout_structure) else state_dict.get('use_dropout', False) + self.activate_output = state_dict.get('activate_output', True) + self.last_layer_dropout = state_dict.get('last_layer_dropout', False) + # Dropout structure should have same length as layer structure, Every digits should be in [0,1), and last digit must be 0. + if self.dropout_structure is None: + self.dropout_structure = parse_dropout_structure(self.layer_structure, self.use_dropout, self.last_layer_dropout) + if shared.opts.print_hypernet_extra: + if self.optional_info is not None: + print(f" INFO:\n {self.optional_info}\n") + print(f" Layer structure: {self.layer_structure}") + print(f" Activation function: {self.activation_func}") + print(f" Weight initialization: {self.weight_init}") + print(f" Layer norm: {self.add_layer_norm}") + print(f" Dropout usage: {self.use_dropout}" ) + print(f" Activate last layer: {self.activate_output}") + print(f" Dropout structure: {self.dropout_structure}") + optimizer_saved_dict = torch.load(self.filename + '.optim', map_location='cpu') if os.path.exists(self.filename + '.optim') else {} + if self.shorthash() == optimizer_saved_dict.get('hash', None): + self.optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None) + else: + self.optimizer_state_dict = None + if self.optimizer_state_dict: + self.optimizer_name = optimizer_saved_dict.get('optimizer_name', 'AdamW') + if shared.opts.print_hypernet_extra: + print("Load existing optimizer from checkpoint") + print(f"Optimizer name is {self.optimizer_name}") + else: + self.optimizer_name = "AdamW" + if shared.opts.print_hypernet_extra: + print("No saved optimizer exists in checkpoint") + for size, sd in state_dict.items(): + if type(size) == int: + self.layers[size] = ( + HypernetworkModule(size, sd[0], self.layer_structure, self.activation_func, self.weight_init, + self.add_layer_norm, self.activate_output, self.dropout_structure), + HypernetworkModule(size, sd[1], self.layer_structure, self.activation_func, self.weight_init, + self.add_layer_norm, self.activate_output, self.dropout_structure), + ) + self.name = state_dict.get('name', self.name) + self.step = state_dict.get('step', 0) + self.sd_checkpoint = state_dict.get('sd_checkpoint', None) + self.sd_checkpoint_name = state_dict.get('sd_checkpoint_name', None) + self.eval() + + def shorthash(self): + sha256 = hashes.sha256(self.filename, f'hypernet/{self.name}') + return sha256[0:10] if sha256 else None + + +def list_hypernetworks(path): + hypernetworks = { + os.path.splitext(os.path.basename(hypernetwork_path))[0]: hypernetwork_path + for hypernetwork_path + in files_cache.list_files( + path, + ext_filter=['.pt'], + recursive=files_cache.not_hidden + ) + } + return hypernetworks + + +def load_hypernetwork(name): + path = shared.hypernetworks.get(name, None) + if path is None: + return None + hypernetwork = Hypernetwork() + try: + hypernetwork.load(path) + except Exception as e: + errors.display(e, f'hypernetwork load: {path}') + return None + return hypernetwork + + +def load_hypernetworks(names, multipliers=None): + already_loaded = {} + for hypernetwork in shared.loaded_hypernetworks: + if hypernetwork.name in names: + already_loaded[hypernetwork.name] = hypernetwork + shared.loaded_hypernetworks.clear() + for i, name in enumerate(names): + hypernetwork = already_loaded.get(name, None) + if hypernetwork is None: + hypernetwork = load_hypernetwork(name) + if hypernetwork is None: + continue + hypernetwork.set_multiplier(multipliers[i] if multipliers else 1.0) + shared.loaded_hypernetworks.append(hypernetwork) + + +def find_closest_hypernetwork_name(search: str): + if not search: + return None + search = search.lower() + applicable = [name for name in shared.hypernetworks if search in name.lower()] + if not applicable: + return None + applicable = sorted(applicable, key=lambda name: len(name)) + return applicable[0] + + +def apply_single_hypernetwork(hypernetwork, context_k, context_v, layer=None): + hypernetwork_layers = (hypernetwork.layers if hypernetwork is not None else {}).get(context_k.shape[2], None) + if hypernetwork_layers is None: + return context_k, context_v + if layer is not None: + layer.hyper_k = hypernetwork_layers[0] + layer.hyper_v = hypernetwork_layers[1] + context_k = devices.cond_cast_unet(hypernetwork_layers[0](devices.cond_cast_float(context_k))) + context_v = devices.cond_cast_unet(hypernetwork_layers[1](devices.cond_cast_float(context_v))) + return context_k, context_v + + +def apply_hypernetworks(hypernetworks, context, layer=None): + context_k = context + context_v = context + for hypernetwork in hypernetworks: + context_k, context_v = apply_single_hypernetwork(hypernetwork, context_k, context_v, layer) + return context_k, context_v + + +def attention_CrossAttention_forward(self, x, context=None, mask=None): + h = self.heads + q = self.to_q(x) + context = default(context, x) + context_k, context_v = apply_hypernetworks(shared.loaded_hypernetworks, context, self) + k = self.to_k(context_k) + v = self.to_v(context_v) + q, k, v = (rearrange(t, 'b n (h d) -> (b h) n d', h=h) for t in (q, k, v)) + sim = einsum('b i d, b j d -> b i j', q, k) * self.scale + if mask is not None: + mask = rearrange(mask, 'b ... -> b (...)') + max_neg_value = -torch.finfo(sim.dtype).max + mask = repeat(mask, 'b j -> (b h) () j', h=h) + sim.masked_fill_(~mask, max_neg_value) + # attention, what we cannot get enough of + attn = sim.softmax(dim=-1) + out = einsum('b i j, b j d -> b i d', attn, v) + out = rearrange(out, '(b h) n d -> b n (h d)', h=h) + return self.to_out(out) + + +def stack_conds(conds): + if len(conds) == 1: + return torch.stack(conds) + # same as in reconstruct_multicond_batch + token_count = max([x.shape[0] for x in conds]) + for i in range(len(conds)): + if conds[i].shape[0] != token_count: + last_vector = conds[i][-1:] + last_vector_repeated = last_vector.repeat([token_count - conds[i].shape[0], 1]) + conds[i] = torch.vstack([conds[i], last_vector_repeated]) + return torch.stack(conds) + + +def statistics(data): + if len(data) < 2: + std = 0 + else: + std = stdev(data) + total_information = f"loss:{mean(data):.3f}" + "\u00B1" + f"({std/ (len(data) ** 0.5):.3f})" + recent_data = data[-32:] + if len(recent_data) < 2: + std = 0 + else: + std = stdev(recent_data) + recent_information = f"recent 32 loss:{mean(recent_data):.3f}" + "\u00B1" + f"({std / (len(recent_data) ** 0.5):.3f})" + return total_information, recent_information + + +def report_statistics(loss_info:dict): + keys = sorted(loss_info.keys(), key=lambda x: sum(loss_info[x]) / len(loss_info[x])) + for key in keys: + try: + print("Loss statistics for file " + key) + info, recent = statistics(list(loss_info[key])) + print(info) + print(recent) + except Exception as e: + print(e) + + +def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, dropout_structure=None): + # Remove illegal characters from name. + name = "".join( x for x in name if (x.isalnum() or x in "._- ")) + assert name, "Name cannot be empty!" + fn = os.path.join(shared.opts.hypernetwork_dir, f"{name}.pt") + if not overwrite_old: + assert not os.path.exists(fn), f"file {fn} already exists" + if type(layer_structure) == str: + layer_structure = [float(x.strip()) for x in layer_structure.split(",")] + if use_dropout and dropout_structure and type(dropout_structure) == str: + dropout_structure = [float(x.strip()) for x in dropout_structure.split(",")] + else: + dropout_structure = [0] * len(layer_structure) + hypernet = modules.hypernetworks.hypernetwork.Hypernetwork( + name=name, + enable_sizes=[int(x) for x in enable_sizes], + layer_structure=layer_structure, + activation_func=activation_func, + weight_init=weight_init, + add_layer_norm=add_layer_norm, + use_dropout=use_dropout, + dropout_structure=dropout_structure + ) + hypernet.save(fn) + shared.reload_hypernetworks() + return name + + +def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_hypernetwork_every, template_filename, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument + # images allows training previews to have infotext. Importing it at the top causes a circular import problem. + from modules import images + + save_hypernetwork_every = save_hypernetwork_every or 0 + create_image_every = create_image_every or 0 + template_file = textual_inversion.textual_inversion_templates.get(template_filename, None) + textual_inversion.validate_train_inputs(hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_hypernetwork_every, create_image_every, name="hypernetwork") + template_file = template_file.path + + path = shared.hypernetworks.get(hypernetwork_name, None) + hypernetwork = Hypernetwork() + hypernetwork.load(path) + shared.loaded_hypernetworks = [hypernetwork] + + shared.state.job = "train" + shared.state.textinfo = "Initializing hypernetwork training..." + shared.state.job_count = steps + + hypernetwork_name = hypernetwork_name.rsplit('(', 1)[0] + filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt') + + log_directory = os.path.join(log_directory, datetime.datetime.now().strftime("%Y-%m-%d"), hypernetwork_name) + unload = shared.opts.unload_models_when_training + + if save_hypernetwork_every > 0: + hypernetwork_dir = os.path.join(log_directory, "hypernetworks") + os.makedirs(hypernetwork_dir, exist_ok=True) + else: + hypernetwork_dir = None + + if create_image_every > 0: + images_dir = os.path.join(log_directory, "images") + os.makedirs(images_dir, exist_ok=True) + else: + images_dir = None + + checkpoint = sd_models.select_checkpoint() + + initial_step = hypernetwork.step or 0 + if initial_step >= steps: + shared.state.textinfo = "Model has already been trained beyond specified max steps" + return hypernetwork, filename + + scheduler = LearnRateScheduler(learn_rate, steps, initial_step) + + clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else None + if clip_grad: + clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False) + + if shared.opts.training_enable_tensorboard: + tensorboard_writer = textual_inversion.tensorboard_setup(log_directory) + + # dataset loading may take a while, so input validations and early returns should be done before this + shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..." + + pin_memory = shared.opts.pin_memory + + ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=hypernetwork_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, include_cond=True, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight) + + if shared.opts.save_training_settings_to_txt: + saved_params = dict( + model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds), + **{field: getattr(hypernetwork, field) for field in ['layer_structure', 'activation_func', 'weight_init', 'add_layer_norm', 'use_dropout', ]} + ) + ti_logging.save_settings_to_file(log_directory, {**saved_params, **locals()}) + + latent_sampling_method = ds.latent_sampling_method + + dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory) + + old_parallel_processing_allowed = shared.parallel_processing_allowed + + if unload: + shared.parallel_processing_allowed = False + shared.sd_model.cond_stage_model.to(devices.cpu) + shared.sd_model.first_stage_model.to(devices.cpu) + + weights = hypernetwork.weights() + hypernetwork.train() + + # Here we use optimizer from saved HN, or we can specify as UI option. + if hypernetwork.optimizer_name in optimizer_dict: + optimizer = optimizer_dict[hypernetwork.optimizer_name](params=weights, lr=scheduler.learn_rate) + optimizer_name = hypernetwork.optimizer_name + else: + print(f"Optimizer type {hypernetwork.optimizer_name} is not defined!") + optimizer = torch.optim.AdamW(params=weights, lr=scheduler.learn_rate) + optimizer_name = 'AdamW' + + if hypernetwork.optimizer_state_dict: # This line must be changed if Optimizer type can be different from saved optimizer. + try: + optimizer.load_state_dict(hypernetwork.optimizer_state_dict) + except RuntimeError as e: + print("Cannot resume from saved optimizer!") + print(e) + + scaler = torch.cuda.amp.GradScaler() + + batch_size = ds.batch_size + gradient_step = ds.gradient_step + # n steps = batch_size * gradient_step * n image processed + steps_per_epoch = len(ds) // batch_size // gradient_step + max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step + loss_step = 0 + _loss_step = 0 #internal + # size = len(ds.indexes) + # loss_dict = defaultdict(lambda : deque(maxlen = 1024)) + loss_logging = deque(maxlen=len(ds) * 3) # this should be configurable parameter, this is 3 * epoch(dataset size) + # losses = torch.zeros((size,)) + # previous_mean_losses = [0] + # previous_mean_loss = 0 + # print("Mean loss of {} elements".format(size)) + + _steps_without_grad = 0 + + last_saved_file = "" + last_saved_image = "" + forced_filename = "" + + pbar = tqdm.tqdm(total=steps - initial_step) + try: + sd_hijack_checkpoint.add() + + for _i in range((steps-initial_step) * gradient_step): + if scheduler.finished: + break + if shared.state.interrupted: + break + for j, batch in enumerate(dl): + # works as a drop_last=True for gradient accumulation + if j == max_steps_per_epoch: + break + scheduler.apply(optimizer, hypernetwork.step) + if scheduler.finished: + break + if shared.state.interrupted: + break + + if clip_grad: + clip_grad_sched.step(hypernetwork.step) + + with devices.autocast(): + x = batch.latent_sample.to(devices.device, non_blocking=pin_memory) + if use_weight: + w = batch.weight.to(devices.device, non_blocking=pin_memory) + if tag_drop_out != 0 or shuffle_tags: + shared.sd_model.cond_stage_model.to(devices.device) + c = shared.sd_model.cond_stage_model(batch.cond_text).to(devices.device, non_blocking=pin_memory) + shared.sd_model.cond_stage_model.to(devices.cpu) + else: + c = stack_conds(batch.cond).to(devices.device, non_blocking=pin_memory) + if use_weight: + loss = shared.sd_model.weighted_forward(x, c, w)[0] / gradient_step + del w + else: + loss = shared.sd_model.forward(x, c)[0] / gradient_step + del x + del c + _loss_step += loss.item() + + scaler.scale(loss).backward() + # go back until we reach gradient accumulation steps + if (j + 1) % gradient_step != 0: + continue + loss_logging.append(_loss_step) + if clip_grad: + clip_grad(weights, clip_grad_sched.learn_rate) + + scaler.step(optimizer) + scaler.update() + hypernetwork.step += 1 + pbar.update() + optimizer.zero_grad(set_to_none=True) + loss_step = _loss_step + _loss_step = 0 + steps_done = hypernetwork.step + 1 + epoch_num = hypernetwork.step // steps_per_epoch + epoch_step = hypernetwork.step % steps_per_epoch + + description = f"Training hypernetwork [Epoch {epoch_num}: {epoch_step+1}/{steps_per_epoch}]loss: {loss_step:.7f}" + pbar.set_description(description) + if hypernetwork_dir is not None and steps_done % save_hypernetwork_every == 0: + # Before saving, change name to match current checkpoint. + hypernetwork_name_every = f'{hypernetwork_name}-{steps_done}' + last_saved_file = os.path.join(hypernetwork_dir, f'{hypernetwork_name_every}.pt') + hypernetwork.optimizer_name = optimizer_name + if shared.opts.save_optimizer_state: + hypernetwork.optimizer_state_dict = optimizer.state_dict() + save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, last_saved_file) + hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory. + + + + if shared.opts.training_enable_tensorboard: + epoch_num = hypernetwork.step // len(ds) + epoch_step = hypernetwork.step - (epoch_num * len(ds)) + 1 + mean_loss = sum(loss_logging) / len(loss_logging) + textual_inversion.tensorboard_add(tensorboard_writer, loss=mean_loss, global_step=hypernetwork.step, step=epoch_step, learn_rate=scheduler.learn_rate, epoch_num=epoch_num) + + textual_inversion.write_loss(log_directory, "hypernetwork_loss.csv", hypernetwork.step, steps_per_epoch, { + "loss": f"{loss_step:.7f}", + "learn_rate": scheduler.learn_rate + }) + + if images_dir is not None and steps_done % create_image_every == 0: + forced_filename = f'{hypernetwork_name}-{steps_done}' + last_saved_image = os.path.join(images_dir, forced_filename) + hypernetwork.eval() + rng_state = torch.get_rng_state() + cuda_rng_state = None + cuda_rng_state = torch.cuda.get_rng_state_all() + shared.sd_model.cond_stage_model.to(devices.device) + shared.sd_model.first_stage_model.to(devices.device) + + p = processing.StableDiffusionProcessingTxt2Img( + sd_model=shared.sd_model, + do_not_save_grid=True, + do_not_save_samples=True, + ) + + p.disable_extra_networks = True + + if preview_from_txt2img: + p.prompt = preview_prompt + p.negative_prompt = preview_negative_prompt + p.steps = preview_steps + p.sampler_name = processing.get_sampler_name(preview_sampler_index) + p.cfg_scale = preview_cfg_scale + p.seed = preview_seed + p.width = preview_width + p.height = preview_height + else: + p.prompt = batch.cond_text[0] + p.steps = 20 + p.width = training_width + p.height = training_height + + preview_text = p.prompt + + processed = processing.process_images(p) + image = processed.images[0] if len(processed.images) > 0 else None + + if unload: + shared.sd_model.cond_stage_model.to(devices.cpu) + shared.sd_model.first_stage_model.to(devices.cpu) + torch.set_rng_state(rng_state) + torch.cuda.set_rng_state_all(cuda_rng_state) + hypernetwork.train() + if image is not None: + shared.state.assign_current_image(image) + if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images: + textual_inversion.tensorboard_add_image(tensorboard_writer, + f"Validation at epoch {epoch_num}", image, + hypernetwork.step) + last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) + last_saved_image += f", prompt: {preview_text}" + + shared.state.job_no = hypernetwork.step + + shared.state.textinfo = f""" +

+Loss: {loss_step:.7f}
+Step: {steps_done}
+Last prompt: {html.escape(batch.cond_text[0])}
+Last saved hypernetwork: {html.escape(last_saved_file)}
+Last saved image: {html.escape(last_saved_image)}
+

+""" + except Exception as e: + errors.display(e, 'hypernetwork train') + finally: + pbar.leave = False + pbar.close() + hypernetwork.eval() + #report_statistics(loss_dict) + sd_hijack_checkpoint.remove() + + + + filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt') + hypernetwork.optimizer_name = optimizer_name + if shared.opts.save_optimizer_state: + hypernetwork.optimizer_state_dict = optimizer.state_dict() + save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename) + + del optimizer + hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory. + shared.sd_model.cond_stage_model.to(devices.device) + shared.sd_model.first_stage_model.to(devices.device) + shared.parallel_processing_allowed = old_parallel_processing_allowed + + return hypernetwork, filename + +def save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename): + old_hypernetwork_name = hypernetwork.name + old_sd_checkpoint = hypernetwork.sd_checkpoint if hasattr(hypernetwork, "sd_checkpoint") else None + old_sd_checkpoint_name = hypernetwork.sd_checkpoint_name if hasattr(hypernetwork, "sd_checkpoint_name") else None + try: + hypernetwork.sd_checkpoint = checkpoint.shorthash + hypernetwork.sd_checkpoint_name = checkpoint.model_name + hypernetwork.name = hypernetwork_name + hypernetwork.save(filename) + except Exception: + hypernetwork.sd_checkpoint = old_sd_checkpoint + hypernetwork.sd_checkpoint_name = old_sd_checkpoint_name + hypernetwork.name = old_hypernetwork_name + raise diff --git a/modules/interrogate.py b/modules/interrogate.py index 974c1f9da..cd43653f1 100644 --- a/modules/interrogate.py +++ b/modules/interrogate.py @@ -1,196 +1,196 @@ -import os -import sys -from collections import namedtuple -from pathlib import Path -import re -import torch -import torch.hub # pylint: disable=ungrouped-imports -from PIL import Image -from torchvision import transforms -from torchvision.transforms.functional import InterpolationMode -from modules import devices, paths, shared, lowvram, errors - - -blip_image_eval_size = 384 -clip_model_name = 'ViT-L/14' -Category = namedtuple("Category", ["name", "topn", "items"]) -re_topn = re.compile(r"\.top(\d+)\.") - - -def category_types(): - return [f.stem for f in Path(shared.interrogator.content_dir).glob('*.txt')] - - -def download_default_clip_interrogate_categories(content_dir): - shared.log.info("Downloading CLIP categories...") - tmpdir = f"{content_dir}_tmp" - cat_types = ["artists", "flavors", "mediums", "movements"] - try: - os.makedirs(tmpdir, exist_ok=True) - for category_type in cat_types: - torch.hub.download_url_to_file(f"https://raw.githubusercontent.com/pharmapsychotic/clip-interrogator/main/clip_interrogator/data/{category_type}.txt", os.path.join(tmpdir, f"{category_type}.txt")) - os.rename(tmpdir, content_dir) - except Exception as e: - errors.display(e, "downloading default CLIP interrogate categories") - finally: - if os.path.exists(tmpdir): - os.removedirs(tmpdir) - - -class InterrogateModels: - blip_model = None - clip_model = None - clip_preprocess = None - dtype = None - running_on_cpu = None - - def __init__(self, content_dir): - self.loaded_categories = None - self.skip_categories = [] - self.content_dir = content_dir - self.running_on_cpu = devices.device_interrogate == torch.device("cpu") - - def categories(self): - if not os.path.exists(self.content_dir): - download_default_clip_interrogate_categories(self.content_dir) - if self.loaded_categories is not None and self.skip_categories == shared.opts.interrogate_clip_skip_categories: - return self.loaded_categories - self.loaded_categories = [] - - if os.path.exists(self.content_dir): - self.skip_categories = shared.opts.interrogate_clip_skip_categories - cat_types = [] - for filename in Path(self.content_dir).glob('*.txt'): - cat_types.append(filename.stem) - if filename.stem in self.skip_categories: - continue - m = re_topn.search(filename.stem) - topn = 1 if m is None else int(m.group(1)) - with open(filename, "r", encoding="utf8") as file: - lines = [x.strip() for x in file.readlines()] - self.loaded_categories.append(Category(name=filename.stem, topn=topn, items=lines)) - return self.loaded_categories - - def create_fake_fairscale(self): - class FakeFairscale: - def checkpoint_wrapper(self): - pass - sys.modules["fairscale.nn.checkpoint.checkpoint_activations"] = FakeFairscale - - def load_blip_model(self): - self.create_fake_fairscale() - import models.blip # pylint: disable=no-name-in-module - import modules.modelloader as modelloader - model_path = os.path.join(paths.models_path, "BLIP") - download_name='model_base_caption_capfilt_large.pth', - shared.log.debug(f'Model interrogate load: type=BLiP model={download_name} path={model_path}') - files = modelloader.load_models( - model_path=model_path, - model_url='https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth', - ext_filter=[".pth"], - download_name=download_name, - ) - blip_model = models.blip.blip_decoder(pretrained=files[0], image_size=blip_image_eval_size, vit='base', med_config=os.path.join(paths.paths["BLIP"], "configs", "med_config.json")) # pylint: disable=c-extension-no-member - blip_model.eval() - - return blip_model - - def load_clip_model(self): - shared.log.debug(f'Model interrogate load: type=CLiP model={clip_model_name} path={shared.opts.clip_models_path}') - import clip - if self.running_on_cpu: - model, preprocess = clip.load(clip_model_name, device="cpu", download_root=shared.opts.clip_models_path) - else: - model, preprocess = clip.load(clip_model_name, download_root=shared.opts.clip_models_path) - model.eval() - model = model.to(devices.device_interrogate) - return model, preprocess - - def load(self): - if self.blip_model is None: - self.blip_model = self.load_blip_model() - if not shared.opts.no_half and not self.running_on_cpu: - self.blip_model = self.blip_model.half() - self.blip_model = self.blip_model.to(devices.device_interrogate) - if self.clip_model is None: - self.clip_model, self.clip_preprocess = self.load_clip_model() - if not shared.opts.no_half and not self.running_on_cpu: - self.clip_model = self.clip_model.half() - self.clip_model = self.clip_model.to(devices.device_interrogate) - self.dtype = next(self.clip_model.parameters()).dtype - - def send_clip_to_ram(self): - if not shared.opts.interrogate_keep_models_in_memory: - if self.clip_model is not None: - self.clip_model = self.clip_model.to(devices.cpu) - - def send_blip_to_ram(self): - if not shared.opts.interrogate_keep_models_in_memory: - if self.blip_model is not None: - self.blip_model = self.blip_model.to(devices.cpu) - - def unload(self): - self.send_clip_to_ram() - self.send_blip_to_ram() - devices.torch_gc() - - def rank(self, image_features, text_array, top_count=1): - import clip - devices.torch_gc() - if shared.opts.interrogate_clip_dict_limit != 0: - text_array = text_array[0:int(shared.opts.interrogate_clip_dict_limit)] - top_count = min(top_count, len(text_array)) - text_tokens = clip.tokenize(list(text_array), truncate=True).to(devices.device_interrogate) - text_features = self.clip_model.encode_text(text_tokens).type(self.dtype) - text_features /= text_features.norm(dim=-1, keepdim=True) - similarity = torch.zeros((1, len(text_array))).to(devices.device_interrogate) - for i in range(image_features.shape[0]): - similarity += (100.0 * image_features[i].unsqueeze(0) @ text_features.T).softmax(dim=-1) - similarity /= image_features.shape[0] - top_probs, top_labels = similarity.cpu().topk(top_count, dim=-1) - return [(text_array[top_labels[0][i].numpy()], (top_probs[0][i].numpy()*100)) for i in range(top_count)] - - def generate_caption(self, pil_image): - gpu_image = transforms.Compose([ - transforms.Resize((blip_image_eval_size, blip_image_eval_size), interpolation=InterpolationMode.BICUBIC), - transforms.ToTensor(), - transforms.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)) - ])(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) - with devices.inference_context(): - caption = self.blip_model.generate(gpu_image, sample=False, num_beams=shared.opts.interrogate_clip_num_beams, min_length=shared.opts.interrogate_clip_min_length, max_length=shared.opts.interrogate_clip_max_length) - return caption[0] - - def interrogate(self, pil_image): - res = "" - shared.state.begin('interrogate') - try: - if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: - lowvram.send_everything_to_cpu() - devices.torch_gc() - self.load() - if isinstance(pil_image, list): - pil_image = pil_image[0] - if isinstance(pil_image, dict) and 'name' in pil_image: - pil_image = Image.open(pil_image['name']) - pil_image = pil_image.convert("RGB") - caption = self.generate_caption(pil_image) - self.send_blip_to_ram() - devices.torch_gc() - res = caption - clip_image = self.clip_preprocess(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) - with devices.inference_context(), devices.autocast(): - image_features = self.clip_model.encode_image(clip_image).type(self.dtype) - image_features /= image_features.norm(dim=-1, keepdim=True) - for _name, topn, items in self.categories(): - matches = self.rank(image_features, items, top_count=topn) - for match, score in matches: - if shared.opts.interrogate_return_ranks: - res += f", ({match}:{score/100:.3f})" - else: - res += f", {match}" - except Exception as e: - errors.display(e, 'interrogate') - res += "" - self.unload() - shared.state.end() - return res +import os +import sys +from collections import namedtuple +from pathlib import Path +import re +import torch +import torch.hub # pylint: disable=ungrouped-imports +from PIL import Image +from torchvision import transforms +from torchvision.transforms.functional import InterpolationMode +from modules import devices, paths, shared, lowvram, errors + + +blip_image_eval_size = 384 +clip_model_name = 'ViT-L/14' +Category = namedtuple("Category", ["name", "topn", "items"]) +re_topn = re.compile(r"\.top(\d+)\.") + + +def category_types(): + return [f.stem for f in Path(shared.interrogator.content_dir).glob('*.txt')] + + +def download_default_clip_interrogate_categories(content_dir): + shared.log.info("Downloading CLIP categories...") + tmpdir = f"{content_dir}_tmp" + cat_types = ["artists", "flavors", "mediums", "movements"] + try: + os.makedirs(tmpdir, exist_ok=True) + for category_type in cat_types: + torch.hub.download_url_to_file(f"https://raw.githubusercontent.com/pharmapsychotic/clip-interrogator/main/clip_interrogator/data/{category_type}.txt", os.path.join(tmpdir, f"{category_type}.txt")) + os.rename(tmpdir, content_dir) + except Exception as e: + errors.display(e, "downloading default CLIP interrogate categories") + finally: + if os.path.exists(tmpdir): + os.removedirs(tmpdir) + + +class InterrogateModels: + blip_model = None + clip_model = None + clip_preprocess = None + dtype = None + running_on_cpu = None + + def __init__(self, content_dir): + self.loaded_categories = None + self.skip_categories = [] + self.content_dir = content_dir + self.running_on_cpu = devices.device_interrogate == torch.device("cpu") + + def categories(self): + if not os.path.exists(self.content_dir): + download_default_clip_interrogate_categories(self.content_dir) + if self.loaded_categories is not None and self.skip_categories == shared.opts.interrogate_clip_skip_categories: + return self.loaded_categories + self.loaded_categories = [] + + if os.path.exists(self.content_dir): + self.skip_categories = shared.opts.interrogate_clip_skip_categories + cat_types = [] + for filename in Path(self.content_dir).glob('*.txt'): + cat_types.append(filename.stem) + if filename.stem in self.skip_categories: + continue + m = re_topn.search(filename.stem) + topn = 1 if m is None else int(m.group(1)) + with open(filename, "r", encoding="utf8") as file: + lines = [x.strip() for x in file.readlines()] + self.loaded_categories.append(Category(name=filename.stem, topn=topn, items=lines)) + return self.loaded_categories + + def create_fake_fairscale(self): + class FakeFairscale: + def checkpoint_wrapper(self): + pass + sys.modules["fairscale.nn.checkpoint.checkpoint_activations"] = FakeFairscale + + def load_blip_model(self): + self.create_fake_fairscale() + import models.blip # pylint: disable=no-name-in-module + import modules.modelloader as modelloader + model_path = os.path.join(paths.models_path, "BLIP") + download_name='model_base_caption_capfilt_large.pth', + shared.log.debug(f'Model interrogate load: type=BLiP model={download_name} path={model_path}') + files = modelloader.load_models( + model_path=model_path, + model_url='https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth', + ext_filter=[".pth"], + download_name=download_name, + ) + blip_model = models.blip.blip_decoder(pretrained=files[0], image_size=blip_image_eval_size, vit='base', med_config=os.path.join(paths.paths["BLIP"], "configs", "med_config.json")) # pylint: disable=c-extension-no-member + blip_model.eval() + + return blip_model + + def load_clip_model(self): + shared.log.debug(f'Model interrogate load: type=CLiP model={clip_model_name} path={shared.opts.clip_models_path}') + import clip + if self.running_on_cpu: + model, preprocess = clip.load(clip_model_name, device="cpu", download_root=shared.opts.clip_models_path) + else: + model, preprocess = clip.load(clip_model_name, download_root=shared.opts.clip_models_path) + model.eval() + model = model.to(devices.device_interrogate) + return model, preprocess + + def load(self): + if self.blip_model is None: + self.blip_model = self.load_blip_model() + if not shared.opts.no_half and not self.running_on_cpu: + self.blip_model = self.blip_model.half() + self.blip_model = self.blip_model.to(devices.device_interrogate) + if self.clip_model is None: + self.clip_model, self.clip_preprocess = self.load_clip_model() + if not shared.opts.no_half and not self.running_on_cpu: + self.clip_model = self.clip_model.half() + self.clip_model = self.clip_model.to(devices.device_interrogate) + self.dtype = next(self.clip_model.parameters()).dtype + + def send_clip_to_ram(self): + if not shared.opts.interrogate_keep_models_in_memory: + if self.clip_model is not None: + self.clip_model = self.clip_model.to(devices.cpu) + + def send_blip_to_ram(self): + if not shared.opts.interrogate_keep_models_in_memory: + if self.blip_model is not None: + self.blip_model = self.blip_model.to(devices.cpu) + + def unload(self): + self.send_clip_to_ram() + self.send_blip_to_ram() + devices.torch_gc() + + def rank(self, image_features, text_array, top_count=1): + import clip + devices.torch_gc() + if shared.opts.interrogate_clip_dict_limit != 0: + text_array = text_array[0:int(shared.opts.interrogate_clip_dict_limit)] + top_count = min(top_count, len(text_array)) + text_tokens = clip.tokenize(list(text_array), truncate=True).to(devices.device_interrogate) + text_features = self.clip_model.encode_text(text_tokens).type(self.dtype) + text_features /= text_features.norm(dim=-1, keepdim=True) + similarity = torch.zeros((1, len(text_array))).to(devices.device_interrogate) + for i in range(image_features.shape[0]): + similarity += (100.0 * image_features[i].unsqueeze(0) @ text_features.T).softmax(dim=-1) + similarity /= image_features.shape[0] + top_probs, top_labels = similarity.cpu().topk(top_count, dim=-1) + return [(text_array[top_labels[0][i].numpy()], (top_probs[0][i].numpy()*100)) for i in range(top_count)] + + def generate_caption(self, pil_image): + gpu_image = transforms.Compose([ + transforms.Resize((blip_image_eval_size, blip_image_eval_size), interpolation=InterpolationMode.BICUBIC), + transforms.ToTensor(), + transforms.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)) + ])(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) + with devices.inference_context(): + caption = self.blip_model.generate(gpu_image, sample=False, num_beams=shared.opts.interrogate_clip_num_beams, min_length=shared.opts.interrogate_clip_min_length, max_length=shared.opts.interrogate_clip_max_length) + return caption[0] + + def interrogate(self, pil_image): + res = "" + shared.state.begin('interrogate') + try: + if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: + lowvram.send_everything_to_cpu() + devices.torch_gc() + self.load() + if isinstance(pil_image, list): + pil_image = pil_image[0] + if isinstance(pil_image, dict) and 'name' in pil_image: + pil_image = Image.open(pil_image['name']) + pil_image = pil_image.convert("RGB") + caption = self.generate_caption(pil_image) + self.send_blip_to_ram() + devices.torch_gc() + res = caption + clip_image = self.clip_preprocess(pil_image).unsqueeze(0).type(self.dtype).to(devices.device_interrogate) + with devices.inference_context(), devices.autocast(): + image_features = self.clip_model.encode_image(clip_image).type(self.dtype) + image_features /= image_features.norm(dim=-1, keepdim=True) + for _name, topn, items in self.categories(): + matches = self.rank(image_features, items, top_count=topn) + for match, score in matches: + if shared.opts.interrogate_return_ranks: + res += f", ({match}:{score/100:.3f})" + else: + res += f", {match}" + except Exception as e: + errors.display(e, 'interrogate') + res += "" + self.unload() + shared.state.end() + return res diff --git a/modules/sd_models.py b/modules/sd_models.py index cb4609907..a04ff0b0f 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1,1342 +1,1342 @@ -import re -import io -import sys -import json -import time -import copy -import logging -import contextlib -import collections -import os.path -from os import mkdir -from urllib import request -from enum import Enum -from rich import progress # pylint: disable=redefined-builtin -import torch -import safetensors.torch -import diffusers -from omegaconf import OmegaConf -import tomesd -from transformers import logging as transformers_logging -import ldm.modules.midas as midas -from ldm.util import instantiate_from_config -from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_inpainting -from modules.timer import Timer -from modules.memstats import memory_stats -from modules.paths import models_path, script_path -from modules.modeldata import model_data - - -transformers_logging.set_verbosity_error() -model_dir = "Stable-diffusion" -model_path = os.path.abspath(os.path.join(paths.models_path, model_dir)) -checkpoints_list = {} -checkpoint_aliases = {} -checkpoints_loaded = collections.OrderedDict() -sd_metadata_file = os.path.join(paths.data_path, "metadata.json") -sd_metadata = None -sd_metadata_pending = 0 -sd_metadata_timer = 0 - - -class CheckpointInfo: - def __init__(self, filename): - self.name = None - self.hash = None - self.filename = filename - self.type = '' - relname = filename - app_path = os.path.abspath(script_path) - - def rel(fn, path): - try: - return os.path.relpath(fn, path) - except Exception: - return fn - - if relname.startswith('..'): - relname = os.path.abspath(relname) - if relname.startswith(shared.opts.ckpt_dir): - relname = rel(filename, shared.opts.ckpt_dir) - elif relname.startswith(shared.opts.diffusers_dir): - relname = rel(filename, shared.opts.diffusers_dir) - elif relname.startswith(model_path): - relname = rel(filename, model_path) - elif relname.startswith(script_path): - relname = rel(filename, script_path) - elif relname.startswith(app_path): - relname = rel(filename, app_path) - else: - relname = os.path.abspath(relname) - relname, ext = os.path.splitext(relname) - ext = ext.lower()[1:] - - if os.path.isfile(filename): # ckpt or safetensor - self.name = relname - self.filename = filename - self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{relname}") - self.type = ext - # self.model_name = os.path.splitext(name.replace("/", "_").replace("\\", "_"))[0] - else: # maybe a diffuser - repo = [r for r in modelloader.diffuser_repos if filename == r['name']] - if len(repo) == 0: - self.name = relname - self.filename = filename - self.sha256 = None - self.type = 'unknown' - else: - self.name = os.path.join(os.path.basename(shared.opts.diffusers_dir), repo[0]['name']) - self.filename = repo[0]['path'] - self.sha256 = repo[0]['hash'] - self.type = 'diffusers' - - self.shorthash = self.sha256[0:10] if self.sha256 else None - self.title = self.name if self.shorthash is None else f'{self.name} [{self.shorthash}]' - self.path = self.filename - self.model_name = os.path.basename(self.name) - self.metadata = read_metadata_from_safetensors(filename) - # shared.log.debug(f'Checkpoint: type={self.type} name={self.name} filename={self.filename} hash={self.shorthash} title={self.title}') - - def register(self): - checkpoints_list[self.title] = self - for i in [self.name, self.filename, self.shorthash, self.title]: - if i is not None: - checkpoint_aliases[i] = self - - def calculate_shorthash(self): - self.sha256 = hashes.sha256(self.filename, f"checkpoint/{self.name}") - if self.sha256 is None: - return None - self.shorthash = self.sha256[0:10] - checkpoints_list.pop(self.title) - self.title = f'{self.name} [{self.shorthash}]' - self.register() - return self.shorthash - - -class NoWatermark: - def apply_watermark(self, img): - return img - - -def setup_model(): - if not os.path.exists(model_path): - os.makedirs(model_path, exist_ok=True) - list_models() - enable_midas_autodownload() - - -def checkpoint_tiles(use_short=False): # pylint: disable=unused-argument - def convert(name): - return int(name) if name.isdigit() else name.lower() - def alphanumeric_key(key): - return [convert(c) for c in re.split('([0-9]+)', key)] - return sorted([x.title for x in checkpoints_list.values()], key=alphanumeric_key) - - -def list_models(): - t0 = time.time() - global checkpoints_list # pylint: disable=global-statement - checkpoints_list.clear() - checkpoint_aliases.clear() - if shared.opts.sd_disable_ckpt or shared.backend == shared.Backend.DIFFUSERS: - ext_filter = [".safetensors"] - else: - ext_filter = [".ckpt", ".safetensors"] - model_list = list(modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) - if shared.backend == shared.Backend.DIFFUSERS: - model_list += modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir, clear=True) - for filename in sorted(model_list, key=str.lower): - checkpoint_info = CheckpointInfo(filename) - if checkpoint_info.name is not None: - checkpoint_info.register() - if shared.cmd_opts.ckpt is not None: - if not os.path.exists(shared.cmd_opts.ckpt) and shared.backend == shared.Backend.ORIGINAL: - if shared.cmd_opts.ckpt.lower() != "none": - shared.log.warning(f"Requested checkpoint not found: {shared.cmd_opts.ckpt}") - else: - checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt) - if checkpoint_info.name is not None: - checkpoint_info.register() - shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title - elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None: - shared.log.warning(f"Checkpoint not found: {shared.cmd_opts.ckpt}") - shared.log.info(f'Available models: path="{shared.opts.ckpt_dir}" items={len(checkpoints_list)} time={time.time()-t0:.2f}') - - checkpoints_list = dict(sorted(checkpoints_list.items(), key=lambda cp: cp[1].filename)) - """ - if len(checkpoints_list) == 0: - if not shared.cmd_opts.no_download: - key = input('Download the default model? (y/N) ') - if key.lower().startswith('y'): - if shared.backend == shared.Backend.ORIGINAL: - model_url = "https://huggingface.co/runwayml/stable-diffusion-v1-5/resolve/main/v1-5-pruned-emaonly.safetensors" - shared.opts.data['sd_model_checkpoint'] = "v1-5-pruned-emaonly.safetensors" - model_list = modelloader.load_models(model_path=model_path, model_url=model_url, command_path=shared.opts.ckpt_dir, ext_filter=[".ckpt", ".safetensors"], download_name="v1-5-pruned-emaonly.safetensors", ext_blacklist=[".vae.ckpt", ".vae.safetensors"]) - else: - default_model_id = "runwayml/stable-diffusion-v1-5" - modelloader.download_diffusers_model(default_model_id, shared.opts.diffusers_dir) - model_list = modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir) - - for filename in sorted(model_list, key=str.lower): - checkpoint_info = CheckpointInfo(filename) - if checkpoint_info.name is not None: - checkpoint_info.register() - """ - -def update_model_hashes(): - txt = [] - lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.hash is None] - # shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') - for ckpt in lst: - ckpt.hash = model_hash(ckpt.filename) - # txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') - # txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') - lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.sha256 is None or ckpt.shorthash is None] - shared.log.info(f'Models list: hash missing={len(lst)} total={len(checkpoints_list)}') - for ckpt in lst: - ckpt.sha256 = hashes.sha256(ckpt.filename, f"checkpoint/{ckpt.name}") - ckpt.shorthash = ckpt.sha256[0:10] if ckpt.sha256 is not None else None - if ckpt.sha256 is not None: - txt.append(f'Calculated full hash: {ckpt.title} {ckpt.shorthash}') - else: - txt.append(f'Skipped hash calculation: {ckpt.title}') - txt.append(f'Updated hashes for {len(lst)} out of {len(checkpoints_list)} models') - txt = '
'.join(txt) - return txt - - -def get_closet_checkpoint_match(search_string): - checkpoint_info = checkpoint_aliases.get(search_string, None) - if checkpoint_info is not None: - return checkpoint_info - found = sorted([info for info in checkpoints_list.values() if search_string in info.title], key=lambda x: len(x.title)) - if found: - return found[0] - found = sorted([info for info in checkpoints_list.values() if search_string.split(' ')[0] in info.title], key=lambda x: len(x.title)) - if found: - return found[0] - return None - - -def model_hash(filename): - """old hash that only looks at a small part of the file and is prone to collisions""" - try: - with open(filename, "rb") as file: - import hashlib - # t0 = time.time() - m = hashlib.sha256() - file.seek(0x100000) - m.update(file.read(0x10000)) - shorthash = m.hexdigest()[0:8] - # t1 = time.time() - # shared.log.debug(f'Calculating short hash: {filename} hash={shorthash} time={(t1-t0):.2f}') - return shorthash - except FileNotFoundError: - return 'NOFILE' - except Exception: - return 'NOHASH' - - -def select_checkpoint(op='model'): - if op == 'dict': - model_checkpoint = shared.opts.sd_model_dict - elif op == 'refiner': - model_checkpoint = shared.opts.data.get('sd_model_refiner', None) - else: - model_checkpoint = shared.opts.sd_model_checkpoint - if model_checkpoint is None or model_checkpoint == 'None': - return None - checkpoint_info = get_closet_checkpoint_match(model_checkpoint) - if checkpoint_info is not None: - shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"') - return checkpoint_info - if len(checkpoints_list) == 0 and not shared.cmd_opts.no_download: - shared.log.warning("Cannot generate without a checkpoint") - shared.log.info("Set system paths to use existing folders in a different location") - shared.log.info("Or use --ckpt to force using existing checkpoint") - return None - checkpoint_info = next(iter(checkpoints_list.values())) - if model_checkpoint is not None: - if model_checkpoint != 'model.ckpt' and model_checkpoint != 'runwayml/stable-diffusion-v1-5': - shared.log.warning(f"Selected checkpoint not found: {model_checkpoint}") - else: - shared.log.info("Selecting first available checkpoint") - # shared.log.warning(f"Loading fallback checkpoint: {checkpoint_info.title}") - shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title - shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"') - return checkpoint_info - - -checkpoint_dict_replacements = { - 'cond_stage_model.transformer.embeddings.': 'cond_stage_model.transformer.text_model.embeddings.', - 'cond_stage_model.transformer.encoder.': 'cond_stage_model.transformer.text_model.encoder.', - 'cond_stage_model.transformer.final_layer_norm.': 'cond_stage_model.transformer.text_model.final_layer_norm.', -} - - -def transform_checkpoint_dict_key(k): - for text, replacement in checkpoint_dict_replacements.items(): - if k.startswith(text): - k = replacement + k[len(text):] - return k - - -def get_state_dict_from_checkpoint(pl_sd): - pl_sd = pl_sd.pop("state_dict", pl_sd) - pl_sd.pop("state_dict", None) - sd = {} - for k, v in pl_sd.items(): - new_key = transform_checkpoint_dict_key(k) - if new_key is not None: - sd[new_key] = v - pl_sd.clear() - pl_sd.update(sd) - return pl_sd - - -def write_metadata(): - global sd_metadata_pending # pylint: disable=global-statement - if sd_metadata_pending == 0: - shared.log.debug(f'Model metadata: file="{sd_metadata_file}" no changes') - return - shared.writefile(sd_metadata, sd_metadata_file) - shared.log.info(f'Model metadata saved: file="{sd_metadata_file}" items={sd_metadata_pending} time={sd_metadata_timer:.2f}') - sd_metadata_pending = 0 - - -def scrub_dict(dict_obj, keys): - for key in list(dict_obj.keys()): - if not isinstance(dict_obj, dict): - continue - if key in keys: - dict_obj.pop(key, None) - elif isinstance(dict_obj[key], dict): - scrub_dict(dict_obj[key], keys) - elif isinstance(dict_obj[key], list): - for item in dict_obj[key]: - scrub_dict(item, keys) - - -def read_metadata_from_safetensors(filename): - global sd_metadata # pylint: disable=global-statement - if sd_metadata is None: - if not os.path.isfile(sd_metadata_file): - sd_metadata = {} - else: - sd_metadata = shared.readfile(sd_metadata_file, lock=True) - res = sd_metadata.get(filename, None) - if res is not None: - return res - if not filename.endswith(".safetensors"): - return {} - if shared.cmd_opts.no_metadata: - return {} - res = {} - try: - t0 = time.time() - with open(filename, mode="rb") as file: - metadata_len = file.read(8) - metadata_len = int.from_bytes(metadata_len, "little") - json_start = file.read(2) - if metadata_len <= 2 or json_start not in (b'{"', b"{'"): - shared.log.error(f"Not a valid safetensors file: {filename}") - json_data = json_start + file.read(metadata_len-2) - json_obj = json.loads(json_data) - for k, v in json_obj.get("__metadata__", {}).items(): - if v.startswith("data:"): - v = 'data' - if k == 'format' and v == 'pt': - continue - large = True if len(v) > 2048 else False - if large and k == 'ss_datasets': - continue - if large and k == 'workflow': - continue - if large and k == 'prompt': - continue - if large and k == 'ss_bucket_info': - continue - if v[0:1] == '{': - try: - v = json.loads(v) - if large and k == 'ss_tag_frequency': - v = { i: len(j) for i, j in v.items() } - if large and k == 'sd_merge_models': - scrub_dict(v, ['sd_merge_recipe']) - except Exception: - pass - res[k] = v - sd_metadata[filename] = res - global sd_metadata_pending # pylint: disable=global-statement - sd_metadata_pending += 1 - t1 = time.time() - global sd_metadata_timer # pylint: disable=global-statement - sd_metadata_timer += (t1 - t0) - except Exception as e: - shared.log.error(f"Error reading metadata from: {filename} {e}") - return res - - -def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unused-argument - if not os.path.isfile(checkpoint_file): - shared.log.error(f"Model is not a file: {checkpoint_file}") - return None - try: - pl_sd = None - with progress.open(checkpoint_file, 'rb', description=f'[cyan]Loading model: [yellow]{checkpoint_file}', auto_refresh=True, console=shared.console) as f: - _, extension = os.path.splitext(checkpoint_file) - if extension.lower() == ".ckpt" and shared.opts.sd_disable_ckpt: - shared.log.warning(f"Checkpoint loading disabled: {checkpoint_file}") - return None - if shared.opts.stream_load: - if extension.lower() == ".safetensors": - # shared.log.debug('Model weights loading: type=safetensors mode=buffered') - buffer = f.read() - pl_sd = safetensors.torch.load(buffer) - else: - # shared.log.debug('Model weights loading: type=checkpoint mode=buffered') - buffer = io.BytesIO(f.read()) - pl_sd = torch.load(buffer, map_location='cpu') - else: - if extension.lower() == ".safetensors": - # shared.log.debug('Model weights loading: type=safetensors mode=mmap') - pl_sd = safetensors.torch.load_file(checkpoint_file, device='cpu') - else: - # shared.log.debug('Model weights loading: type=checkpoint mode=direct') - pl_sd = torch.load(f, map_location='cpu') - sd = get_state_dict_from_checkpoint(pl_sd) - del pl_sd - except Exception as e: - errors.display(e, f'Load model: {checkpoint_file}') - sd = None - return sd - - -def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer): - if not os.path.isfile(checkpoint_info.filename): - return None - if checkpoint_info in checkpoints_loaded: - shared.log.info("Model weights loading: from cache") - return checkpoints_loaded[checkpoint_info] - res = read_state_dict(checkpoint_info.filename) - timer.record("load") - return res - - -def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo, state_dict, timer): - _pipeline, _model_type = detect_pipeline(checkpoint_info.path, 'model') - shared.log.debug(f'Model weights loading: {memory_stats()}') - timer.record("hash") - if model_data.sd_dict == 'None': - shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title - if state_dict is None: - state_dict = get_checkpoint_state_dict(checkpoint_info, timer) - try: - model.load_state_dict(state_dict, strict=False) - except Exception as e: - shared.log.error(f'Error loading model weights: {checkpoint_info.filename}') - shared.log.error(' '.join(str(e).splitlines()[:2])) - return False - del state_dict - timer.record("apply") - if shared.opts.sd_checkpoint_cache > 0: - # cache newly loaded model - checkpoints_loaded[checkpoint_info] = model.state_dict().copy() - if shared.opts.opt_channelslast: - model.to(memory_format=torch.channels_last) - timer.record("channels") - if not shared.opts.no_half: - vae = model.first_stage_model - depth_model = getattr(model, 'depth_model', None) - # with --no-half-vae, remove VAE from model when doing half() to prevent its weights from being converted to float16 - if shared.opts.no_half_vae: - model.first_stage_model = None - # with --upcast-sampling, don't convert the depth model weights to float16 - if shared.opts.upcast_sampling and depth_model: - model.depth_model = None - model.half() - model.first_stage_model = vae - if depth_model: - model.depth_model = depth_model - if shared.opts.cuda_cast_unet: - devices.dtype_unet = model.model.diffusion_model.dtype - else: - model.model.diffusion_model.to(devices.dtype_unet) - model.first_stage_model.to(devices.dtype_vae) - # clean up cache if limit is reached - while len(checkpoints_loaded) > shared.opts.sd_checkpoint_cache: - checkpoints_loaded.popitem(last=False) - model.sd_model_hash = checkpoint_info.calculate_shorthash() - model.sd_model_checkpoint = checkpoint_info.filename - model.sd_checkpoint_info = checkpoint_info - model.is_sdxl = False # a1111 compatibility item - model.is_sd2 = hasattr(model.cond_stage_model, 'model') # a1111 compatibility item - model.is_sd1 = not hasattr(model.cond_stage_model, 'model') # a1111 compatibility item - model.logvar = model.logvar.to(devices.device) if hasattr(model, 'logvar') else None # fix for training - shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256 - sd_vae.delete_base_vae() - sd_vae.clear_loaded_vae() - vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) - sd_vae.load_vae(model, vae_file, vae_source) - timer.record("vae") - return True - - -def enable_midas_autodownload(): - """ - Gives the ldm.modules.midas.api.load_model function automatic downloading. - - When the 512-depth-ema model, and other future models like it, is loaded, - it calls midas.api.load_model to load the associated midas depth model. - This function applies a wrapper to download the model to the correct - location automatically. - """ - midas_path = os.path.join(paths.models_path, 'midas') - for k, v in midas.api.ISL_PATHS.items(): - file_name = os.path.basename(v) - midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) - midas_urls = { - "dpt_large": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_large-midas-2f21e586.pt", - "dpt_hybrid": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_hybrid-midas-501f0c75.pt", - "midas_v21": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21-f6b98070.pt", - "midas_v21_small": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21_small-70d6b9c8.pt", - } - midas.api.load_model_inner = midas.api.load_model - - def load_model_wrapper(model_type): - path = midas.api.ISL_PATHS[model_type] - if not os.path.exists(path): - if not os.path.exists(midas_path): - mkdir(midas_path) - shared.log.info(f"Downloading midas model weights for {model_type} to {path}") - request.urlretrieve(midas_urls[model_type], path) - shared.log.info(f"{model_type} downloaded") - return midas.api.load_model_inner(model_type) - - midas.api.load_model = load_model_wrapper - - -def repair_config(sd_config): - if "use_ema" not in sd_config.model.params: - sd_config.model.params.use_ema = False - if shared.opts.no_half: - sd_config.model.params.unet_config.params.use_fp16 = False - elif shared.opts.upcast_sampling: - sd_config.model.params.unet_config.params.use_fp16 = True if sys.platform != 'darwin' else False - if getattr(sd_config.model.params.first_stage_config.params.ddconfig, "attn_type", None) == "vanilla-xformers" and not shared.xformers_available: - sd_config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla" - # For UnCLIP-L, override the hardcoded karlo directory - if "noise_aug_config" in sd_config.model.params and "clip_stats_path" in sd_config.model.params.noise_aug_config.params: - karlo_path = os.path.join(paths.models_path, 'karlo') - sd_config.model.params.noise_aug_config.params.clip_stats_path = sd_config.model.params.noise_aug_config.params.clip_stats_path.replace("checkpoints/karlo_models", karlo_path) - - -sd1_clip_weight = 'cond_stage_model.transformer.text_model.embeddings.token_embedding.weight' -sd2_clip_weight = 'cond_stage_model.model.transformer.resblocks.0.attn.in_proj_weight' - - -def change_backend(): - shared.log.info(f'Backend changed: {shared.backend}') - shared.log.warning('Full server restart required to apply all changes') - if shared.backend == shared.Backend.ORIGINAL: - change_from = shared.Backend.DIFFUSERS - else: - change_from = shared.Backend.ORIGINAL - unload_model_weights(change_from=change_from) - checkpoints_loaded.clear() - from modules.sd_samplers import list_samplers - list_samplers(shared.backend) - list_models() - from modules.sd_vae import refresh_vae_list - refresh_vae_list() - - -def detect_pipeline(f: str, op: str = 'model', warning=True): - if not f.endswith('.safetensors'): - return None, None - guess = shared.opts.diffusers_pipeline - warn = shared.log.warning if warning else lambda *args, **kwargs: None - if guess == 'Autodetect': - try: - # guess by size - size = round(os.path.getsize(f) / 1024 / 1024) - if size < 128: - warn(f'Model size smaller than expected: {f} size={size} MB') - elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160 - warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB') - guess = 'VAE' - elif size >= 5351 and size <= 5359: # 5353 - guess = 'Stable Diffusion' # SD v2 - elif size >= 5791 and size <= 5799: # 5795 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as SD-XL refiner model, but attempting to load using backend=original: {op}={f} size={size} MB') - if op == 'model': - warn(f'Model detected as SD-XL refiner model, but attempting to load a base model: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL' - elif (size >= 6611 and size <= 6619) or (size >= 6771 and size <= 6779): # 6617, HassakuXL is 6776 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as SD-XL base model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL' - elif size >= 3361 and size <= 3369: # 3368 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as SD upscale model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Stable Diffusion Upscale' - elif size >= 4891 and size <= 4899: # 4897 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as SD XL inpaint model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL Inpaint' - elif size >= 9791 and size <= 9799: # 9794 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as SD XL instruct pix2pix model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL Instruct' - elif size > 3138 and size < 3142: #3140 - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as Segmind Vega model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL' - else: - guess = 'Stable Diffusion' - # guess by name - """ - if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper(): - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as LCM model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Latent Consistency Model' - """ - if 'PixArt' in f: - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as PixArt Alpha model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'PixArt Alpha' - # switch for specific variant - if guess == 'Stable Diffusion' and 'inpaint' in f.lower(): - guess = 'Stable Diffusion Inpaint' - elif guess == 'Stable Diffusion' and 'instruct' in f.lower(): - guess = 'Stable Diffusion Instruct' - if guess == 'Stable Diffusion XL' and 'inpaint' in f.lower(): - guess = 'Stable Diffusion XL Inpaint' - elif guess == 'Stable Diffusion XL' and 'instruct' in f.lower(): - guess = 'Stable Diffusion XL Instruct' - # get actual pipeline - pipeline = shared_items.get_pipelines().get(guess, None) - shared.log.info(f'Autodetect: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB') - except Exception as e: - shared.log.error(f'Error detecting diffusers pipeline: model={f} {e}') - return None, None - else: - try: - size = round(os.path.getsize(f) / 1024 / 1024) - pipeline = shared_items.get_pipelines().get(guess, None) - shared.log.info(f'Diffusers: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB') - except Exception as e: - shared.log.error(f'Error loading diffusers pipeline: model={f} {e}') - - if pipeline is None: - shared.log.warning(f'Autodetect: pipeline not recognized: {guess}: {op}={f} size={size}') - pipeline = diffusers.StableDiffusionPipeline - return pipeline, guess - - -def copy_diffuser_options(new_pipe, orig_pipe): - new_pipe.sd_checkpoint_info = orig_pipe.sd_checkpoint_info - new_pipe.sd_model_checkpoint = orig_pipe.sd_model_checkpoint - new_pipe.embedding_db = getattr(orig_pipe, 'embedding_db', None) - new_pipe.sd_model_hash = getattr(orig_pipe, 'sd_model_hash', None) - new_pipe.has_accelerate = getattr(orig_pipe, 'has_accelerate', False) - new_pipe.is_sdxl = getattr(orig_pipe, 'is_sdxl', False) # a1111 compatibility item - new_pipe.is_sd2 = getattr(orig_pipe, 'is_sd2', False) - new_pipe.is_sd1 = getattr(orig_pipe, 'is_sd1', True) - - -def set_diffuser_options(sd_model, vae = None, op: str = 'model'): - if sd_model is None: - shared.log.warning(f'{op} is not loaded') - return - if (shared.opts.diffusers_model_cpu_offload or shared.cmd_opts.medvram) and (shared.opts.diffusers_seq_cpu_offload or shared.cmd_opts.lowvram): - shared.log.warning(f'Setting {op}: Model CPU offload and Sequential CPU offload are not compatible') - shared.log.debug(f'Setting {op}: disabling model CPU offload') - shared.opts.diffusers_model_cpu_offload=False - shared.cmd_opts.medvram=False - - if hasattr(sd_model, "watermark"): - sd_model.watermark = NoWatermark() - sd_model.has_accelerate = False - if hasattr(sd_model, "vae"): - if vae is not None: - sd_model.vae = vae - if shared.opts.diffusers_vae_upcast != 'default': - if shared.opts.diffusers_vae_upcast == 'true': - sd_model.vae.config.force_upcast = True - else: - sd_model.vae.config.force_upcast = False - if shared.opts.no_half_vae: - devices.dtype_vae = torch.float32 - sd_model.vae.to(devices.dtype_vae) - shared.log.debug(f'Setting {op} VAE: name={sd_vae.loaded_vae_file} upcast={sd_model.vae.config.get("force_upcast", None)}') - if hasattr(sd_model, "enable_model_cpu_offload"): - if (shared.cmd_opts.medvram and devices.backend != "directml") or shared.opts.diffusers_model_cpu_offload: - shared.log.debug(f'Setting {op}: enable model CPU offload') - if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: - shared.opts.diffusers_move_base = False - shared.opts.diffusers_move_unet = False - shared.opts.diffusers_move_refiner = False - shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled') - sd_model.enable_model_cpu_offload() - sd_model.has_accelerate = True - if hasattr(sd_model, "enable_sequential_cpu_offload"): - if shared.cmd_opts.lowvram or shared.opts.diffusers_seq_cpu_offload: - shared.log.debug(f'Setting {op}: enable sequential CPU offload') - if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: - shared.opts.diffusers_move_base = False - shared.opts.diffusers_move_unet = False - shared.opts.diffusers_move_refiner = False - shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled') - sd_model.enable_sequential_cpu_offload(device=devices.device) - sd_model.has_accelerate = True - if hasattr(sd_model, "enable_vae_slicing"): - if shared.cmd_opts.lowvram or shared.opts.diffusers_vae_slicing: - shared.log.debug(f'Setting {op}: enable VAE slicing') - sd_model.enable_vae_slicing() - else: - sd_model.disable_vae_slicing() - if hasattr(sd_model, "enable_vae_tiling"): - if shared.cmd_opts.lowvram or shared.opts.diffusers_vae_tiling: - shared.log.debug(f'Setting {op}: enable VAE tiling') - sd_model.enable_vae_tiling() - else: - sd_model.disable_vae_tiling() - if hasattr(sd_model, "enable_attention_slicing"): - if shared.cmd_opts.lowvram or shared.opts.diffusers_attention_slicing: - shared.log.debug(f'Setting {op}: enable attention slicing') - sd_model.enable_attention_slicing() - else: - sd_model.disable_attention_slicing() - if hasattr(sd_model, "vqvae"): - sd_model.vqvae.to(torch.float32) # vqvae is producing nans in fp16 - if shared.opts.cross_attention_optimization == "xFormers" and hasattr(sd_model, 'enable_xformers_memory_efficient_attention'): - sd_model.enable_xformers_memory_efficient_attention() - if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'fuse_qkv_projections'): - shared.log.debug(f'Setting {op}: enable fused projections') - sd_model.fuse_qkv_projections() - if shared.opts.diffusers_eval: - if hasattr(sd_model, "unet") and hasattr(sd_model.unet, "requires_grad_"): - sd_model.unet.requires_grad_(False) - sd_model.unet.eval() - if hasattr(sd_model, "vae") and hasattr(sd_model.vae, "requires_grad_"): - sd_model.vae.requires_grad_(False) - sd_model.vae.eval() - if hasattr(sd_model, "text_encoder") and hasattr(sd_model.text_encoder, "requires_grad_"): - sd_model.text_encoder.requires_grad_(False) - sd_model.text_encoder.eval() - if shared.opts.diffusers_quantization: - sd_model = sd_models_compile.dynamic_quantization(sd_model) - - if shared.opts.opt_channelslast and hasattr(sd_model, 'unet'): - shared.log.debug(f'Setting {op}: enable channels last') - sd_model.unet.to(memory_format=torch.channels_last) - - -def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): # pylint: disable=unused-argument - import torch # pylint: disable=reimported,redefined-outer-name - if shared.cmd_opts.profile: - import cProfile - pr = cProfile.Profile() - pr.enable() - if timer is None: - timer = Timer() - logging.getLogger("diffusers").setLevel(logging.ERROR) - timer.record("diffusers") - devices.set_cuda_params() - diffusers_load_config = { - "low_cpu_mem_usage": True, - "torch_dtype": devices.dtype, - "safety_checker": None, - "requires_safety_checker": False, - "load_safety_checker": False, - "load_connected_pipeline": True, - # TODO: use_safetensors cant enable for all checkpoints just yet - } - if shared.opts.diffusers_model_load_variant == 'default': - if devices.dtype == torch.float16: - diffusers_load_config['variant'] = 'fp16' - elif shared.opts.diffusers_model_load_variant == 'fp32': - pass - else: - diffusers_load_config['variant'] = shared.opts.diffusers_model_load_variant - - if shared.opts.diffusers_pipeline == 'Custom Diffusers Pipeline' and len(shared.opts.custom_diffusers_pipeline) > 0: - shared.log.debug(f'Diffusers custom pipeline: {shared.opts.custom_diffusers_pipeline}') - diffusers_load_config['custom_pipeline'] = shared.opts.custom_diffusers_pipeline - - # if 'LCM' in checkpoint_info.path: - # diffusers_load_config['custom_pipeline'] = 'latent_consistency_txt2img' - - if shared.opts.data.get('sd_model_checkpoint', '') == 'model.ckpt' or shared.opts.data.get('sd_model_checkpoint', '') == '': - shared.opts.data['sd_model_checkpoint'] = "runwayml/stable-diffusion-v1-5" - - if op == 'model' or op == 'dict': - if (model_data.sd_model is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model - return - else: - if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model - return - - sd_model = None - - try: - if shared.cmd_opts.ckpt is not None and os.path.isdir(shared.cmd_opts.ckpt) and model_data.initial: # initial load - ckpt_basename = os.path.basename(shared.cmd_opts.ckpt) - model_name = modelloader.find_diffuser(ckpt_basename) - if model_name is not None: - shared.log.info(f'Load model {op}: {model_name}') - model_file = modelloader.download_diffusers_model(hub_id=model_name) - try: - shared.log.debug(f'Model load {op} config: {diffusers_load_config}') - sd_model = diffusers.DiffusionPipeline.from_pretrained(model_file, **diffusers_load_config) - except Exception as e: - shared.log.error(f'Failed loading model: {model_file} {e}') - list_models() # rescan for downloaded model - checkpoint_info = CheckpointInfo(model_name) - - checkpoint_info = checkpoint_info or select_checkpoint(op=op) - if checkpoint_info is None: - unload_model_weights(op=op) - return - - vae = None - sd_vae.loaded_vae_file = None - if op == 'model' or op == 'refiner': - vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) - vae = sd_vae.load_vae_diffusers(checkpoint_info.path, vae_file, vae_source) - if vae is not None: - diffusers_load_config["vae"] = vae - - shared.log.debug(f'Diffusers loading: path="{checkpoint_info.path}"') - if os.path.isdir(checkpoint_info.path): - err1 = None - err2 = None - err3 = None - try: # try autopipeline first, best choice but not all pipelines are available - sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - except Exception as e: - err1 = e - # shared.log.error(f'AutoPipeline: {e}') - try: # try diffusion pipeline next second-best choice, works for most non-linked pipelines - if err1 is not None: - sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - except Exception as e: - err2 = e - # shared.log.error(f'DiffusionPipeline: {e}') - try: # try basic pipeline next just in case - if err2 is not None: - sd_model = diffusers.StableDiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - except Exception as e: - err3 = e # ignore last error - shared.log.error(f'StableDiffusionPipeline: {e}') - if err3 is not None: - shared.log.error(f'Failed loading {op}: {checkpoint_info.path} auto={err1} diffusion={err2}') - return - elif os.path.isfile(checkpoint_info.path) and checkpoint_info.path.lower().endswith('.safetensors'): - # diffusers_load_config["local_files_only"] = True - diffusers_load_config["extract_ema"] = shared.opts.diffusers_extract_ema - pipeline, model_type = detect_pipeline(checkpoint_info.path, op) - if pipeline is None: - shared.log.error(f'Diffusers {op} pipeline not initialized: {shared.opts.diffusers_pipeline}') - return - try: - if model_type.startswith('Stable Diffusion'): - diffusers_load_config['force_zeros_for_empty_prompt '] = shared.opts.diffusers_force_zeros - diffusers_load_config['requires_aesthetics_score'] = shared.opts.diffusers_aesthetics_score - if 'inpainting' in checkpoint_info.path.lower(): - diffusers_load_config['config_files'] = { - 'v1': 'configs/v1-inpainting-inference.yaml', - 'v2': 'configs/v2-inference-768-v.yaml', - 'xl': 'configs/sd_xl_base.yaml', - 'xl_refiner': 'configs/sd_xl_refiner.yaml', - } - else: - diffusers_load_config['config_files'] = { - 'v1': 'configs/v1-inference.yaml', - 'v2': 'configs/v2-inference-768-v.yaml', - 'xl': 'configs/sd_xl_base.yaml', - 'xl_refiner': 'configs/sd_xl_refiner.yaml', - } - if hasattr(pipeline, 'from_single_file'): - diffusers_load_config['use_safetensors'] = True - sd_model = pipeline.from_single_file(checkpoint_info.path, **diffusers_load_config) - if sd_model is not None and hasattr(sd_model, 'unet') and hasattr(sd_model.unet, 'config') and 'inpainting' in checkpoint_info.path.lower(): - shared.log.debug('Model patch: type=inpaint') - sd_model.unet.config.in_channels = 9 - elif hasattr(pipeline, 'from_ckpt'): - sd_model = pipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config) - else: - shared.log.error(f'Diffusers {op} cannot load safetensor model: {checkpoint_info.path} {shared.opts.diffusers_pipeline}') - return - if sd_model is not None: - diffusers_load_config.pop('vae', None) - diffusers_load_config.pop('safety_checker', None) - diffusers_load_config.pop('requires_safety_checker', None) - diffusers_load_config.pop('load_safety_checker', None) - diffusers_load_config.pop('config_files', None) - diffusers_load_config.pop('local_files_only', None) - shared.log.debug(f'Setting {op}: pipeline={sd_model.__class__.__name__} config={diffusers_load_config}') # pylint: disable=protected-access - except Exception as e: - shared.log.error(f'Diffusers failed loading: {op}={checkpoint_info.path} pipeline={shared.opts.diffusers_pipeline}/{sd_model.__class__.__name__} {e}') - errors.display(e, f'loading {op}={checkpoint_info.path} pipeline={shared.opts.diffusers_pipeline}/{sd_model.__class__.__name__}') - return - else: - shared.log.error(f'Diffusers cannot load: {op}={checkpoint_info.path}') - return - - if "StableDiffusion" in sd_model.__class__.__name__: - pass # scheduler is created on first use - elif "Kandinsky" in sd_model.__class__.__name__: - sd_model.scheduler.name = 'DDIM' - - set_diffuser_options(sd_model, vae, op) - - base_sent_to_cpu=False - if (shared.opts.cuda_compile and shared.opts.cuda_compile_backend != 'none') or shared.opts.ipex_optimize: - if op == 'refiner' and not getattr(sd_model, 'has_accelerate', False): - gpu_vram = memory_stats().get('gpu', {}) - free_vram = gpu_vram.get('total', 0) - gpu_vram.get('used', 0) - refiner_enough_vram = free_vram >= 7 if "StableDiffusionXL" in sd_model.__class__.__name__ else 3 - if not shared.opts.diffusers_move_base and refiner_enough_vram: - sd_model.to(devices.device) - base_sent_to_cpu=False - else: - if not refiner_enough_vram and not (shared.opts.diffusers_move_base and shared.opts.diffusers_move_refiner): - shared.log.warning(f"Insufficient GPU memory, using system memory as fallback: free={free_vram} GB") - if not shared.opts.shared.opts.diffusers_seq_cpu_offload and not shared.opts.diffusers_model_cpu_offload: - shared.log.debug('Enabled moving base model to CPU') - shared.log.debug('Enabled moving refiner model to CPU') - shared.opts.diffusers_move_base=True - shared.opts.diffusers_move_refiner=True - shared.log.debug('Moving base model to CPU') - if model_data.sd_model is not None: - model_data.sd_model.to(devices.cpu) - devices.torch_gc(force=True) - sd_model.to(devices.device) - base_sent_to_cpu=True - elif not getattr(sd_model, 'has_accelerate', False): - sd_model.to(devices.device) - - sd_models_compile.compile_diffusers(sd_model) - - if sd_model is None: - shared.log.error('Diffuser model not loaded') - return - sd_model.sd_model_hash = checkpoint_info.calculate_shorthash() # pylint: disable=attribute-defined-outside-init - sd_model.sd_checkpoint_info = checkpoint_info # pylint: disable=attribute-defined-outside-init - sd_model.sd_model_checkpoint = checkpoint_info.filename # pylint: disable=attribute-defined-outside-init - sd_model.is_sdxl = False # a1111 compatibility item - sd_model.is_sd2 = hasattr(sd_model, 'cond_stage_model') and hasattr(sd_model.cond_stage_model, 'model') # a1111 compatibility item - sd_model.is_sd1 = not sd_model.is_sd2 # a1111 compatibility item - sd_model.logvar = sd_model.logvar.to(devices.device) if hasattr(sd_model, 'logvar') else None # fix for training - shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256 - if hasattr(sd_model, "set_progress_bar_config"): - sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba') - if op == 'refiner' and shared.opts.diffusers_move_refiner and not getattr(sd_model, 'has_accelerate', False): - shared.log.debug('Moving refiner model to CPU') - sd_model.to(devices.cpu) - elif not getattr(sd_model, 'has_accelerate', False): # In offload modes, accelerate will move models around - sd_model.to(devices.device) - if op == 'refiner' and base_sent_to_cpu: - shared.log.debug('Moving base model back to GPU') - model_data.sd_model.to(devices.device) - except Exception as e: - shared.log.error("Failed to load diffusers model") - errors.display(e, "loading Diffusers model") - - if sd_model is not None: - from modules.textual_inversion import textual_inversion - sd_model.embedding_db = textual_inversion.EmbeddingDatabase() - if op == 'refiner': - model_data.sd_refiner = sd_model - else: - model_data.sd_model = sd_model - sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir) - sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) - - timer.record("load") - devices.torch_gc(force=True) - if shared.cmd_opts.profile: - errors.profile(pr, 'Load') - script_callbacks.model_loaded_callback(sd_model) - shared.log.info(f"Load {op}: time={timer.summary()} native={get_native(sd_model)} {memory_stats()}") - - -class DiffusersTaskType(Enum): - TEXT_2_IMAGE = 1 - IMAGE_2_IMAGE = 2 - INPAINTING = 3 - INSTRUCT = 4 - - -def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType: - if pipe.__class__.__name__ == "StableDiffusionXLInstructPix2PixPipeline": - return DiffusersTaskType.INSTRUCT - elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING.values(): - return DiffusersTaskType.IMAGE_2_IMAGE - elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values(): - return DiffusersTaskType.INPAINTING - else: - return DiffusersTaskType.TEXT_2_IMAGE - - -def switch_diffuser_pipe(pipeline, cls): - try: - new_pipe = None - if isinstance(pipeline, cls): - return pipeline - elif isinstance(pipeline, diffusers.StableDiffusionXLPipeline): - new_pipe = cls( - vae=pipeline.vae, - text_encoder=pipeline.text_encoder, - text_encoder_2=pipeline.text_encoder_2, - tokenizer=pipeline.tokenizer, - tokenizer_2=pipeline.tokenizer_2, - unet=pipeline.unet, - scheduler=pipeline.scheduler, - feature_extractor=getattr(pipeline, 'feature_extractor', None), - ).to(pipeline.device) - elif isinstance(pipeline, diffusers.StableDiffusionPipeline): - new_pipe = cls( - vae=pipeline.vae, - text_encoder=pipeline.text_encoder, - tokenizer=pipeline.tokenizer, - unet=pipeline.unet, - scheduler=pipeline.scheduler, - feature_extractor=getattr(pipeline, 'feature_extractor', None), - requires_safety_checker=False, - safety_checker=None, - ).to(pipeline.device) - else: - shared.log.error(f'Pipeline switch error: {pipeline.__class__.__name__} unrecognized') - return pipeline - if new_pipe is not None: - copy_diffuser_options(new_pipe, pipeline) - shared.log.debug(f'Pipeline switch: from={pipeline.__class__.__name__} to={new_pipe.__class__.__name__}') - return new_pipe - else: - shared.log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls.__name__} empty pipeline') - except Exception as e: - shared.log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls.__name__} {e}') - return pipeline - - -def set_diffuser_pipe(pipe, new_pipe_type): - sd_checkpoint_info = getattr(pipe, "sd_checkpoint_info", None) - sd_model_checkpoint = getattr(pipe, "sd_model_checkpoint", None) - sd_model_hash = getattr(pipe, "sd_model_hash", None) - has_accelerate = getattr(pipe, "has_accelerate", None) - embedding_db = getattr(pipe, "embedding_db", None) - image_encoder = getattr(pipe, "image_encoder", None) - feature_extractor = getattr(pipe, "feature_extractor", None) - - # skip specific pipelines - if pipe.__class__.__name__ == 'StableDiffusionReferencePipeline' or pipe.__class__.__name__ == 'StableDiffusionAdapterPipeline': - return pipe - - try: - if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: - new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: - new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.INPAINTING: - new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) - except Exception as e: # pylint: disable=unused-variable - shared.log.warning(f'Failed to change: type={new_pipe_type} pipeline={pipe.__class__.__name__} {e}') - return pipe - - if pipe.__class__ == new_pipe.__class__: - return pipe - new_pipe.sd_checkpoint_info = sd_checkpoint_info - new_pipe.sd_model_checkpoint = sd_model_checkpoint - new_pipe.sd_model_hash = sd_model_hash - new_pipe.has_accelerate = has_accelerate - new_pipe.embedding_db = embedding_db - new_pipe.image_encoder = image_encoder - new_pipe.feature_extractor = feature_extractor - new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item - new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False) - new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) - shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__}") - pipe = new_pipe - return pipe - - -def get_native(pipe: diffusers.DiffusionPipeline): - if hasattr(pipe, "vae") and hasattr(pipe.vae.config, "sample_size"): - # Stable Diffusion - size = pipe.vae.config.sample_size - elif hasattr(pipe, "movq") and hasattr(pipe.movq.config, "sample_size"): - # Kandinsky - size = pipe.movq.config.sample_size - elif hasattr(pipe, "unet") and hasattr(pipe.unet.config, "sample_size"): - size = pipe.unet.config.sample_size - else: - size = 0 - return size - - -def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): - from modules import lowvram, sd_hijack - checkpoint_info = checkpoint_info or select_checkpoint(op=op) - if checkpoint_info is None: - return - if op == 'model' or op == 'dict': - if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model - return - else: - if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model - return - shared.log.debug(f'Load {op}: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}') - if timer is None: - timer = Timer() - current_checkpoint_info = None - if op == 'model' or op == 'dict': - if model_data.sd_model is not None: - sd_hijack.model_hijack.undo_hijack(model_data.sd_model) - current_checkpoint_info = model_data.sd_model.sd_checkpoint_info - unload_model_weights(op=op) - else: - if model_data.sd_refiner is not None: - sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner) - current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info - unload_model_weights(op=op) - - sd_hijack_inpainting.do_inpainting_hijack() - devices.set_cuda_params() - if already_loaded_state_dict is not None: - state_dict = already_loaded_state_dict - else: - state_dict = get_checkpoint_state_dict(checkpoint_info, timer) - checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info) - if state_dict is None or checkpoint_config is None: - shared.log.error(f"Failed to load checkpooint: {checkpoint_info.filename}") - if current_checkpoint_info is not None: - shared.log.info(f"Restoring previous checkpoint: {current_checkpoint_info.filename}") - load_model(current_checkpoint_info, None) - return - shared.log.debug(f'Model dict loaded: {memory_stats()}') - sd_config = OmegaConf.load(checkpoint_config) - repair_config(sd_config) - timer.record("config") - shared.log.debug(f'Model config loaded: {memory_stats()}') - sd_model = None - stdout = io.StringIO() - if os.environ.get('SD_LDM_DEBUG', None) is not None: - sd_model = instantiate_from_config(sd_config.model) - else: - with contextlib.redirect_stdout(stdout): - """ - try: - clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict - with sd_disable_initialization.DisableInitialization(disable_clip=clip_is_included_into_sd): - sd_model = instantiate_from_config(sd_config.model) - except Exception as e: - shared.log.error(f'LDM: instantiate from config: {e}') - sd_model = instantiate_from_config(sd_config.model) - """ - sd_model = instantiate_from_config(sd_config.model) - for line in stdout.getvalue().splitlines(): - if len(line) > 0: - shared.log.info(f'LDM: {line.strip()}') - shared.log.debug(f"Model created from config: {checkpoint_config}") - sd_model.used_config = checkpoint_config - sd_model.has_accelerate = False - timer.record("create") - ok = load_model_weights(sd_model, checkpoint_info, state_dict, timer) - if not ok: - model_data.sd_model = sd_model - current_checkpoint_info = None - unload_model_weights(op=op) - shared.log.debug(f'Model weights unloaded: {memory_stats()} op={op}') - if op == 'refiner': - # shared.opts.data['sd_model_refiner'] = 'None' - shared.opts.sd_model_refiner = 'None' - return - else: - shared.log.debug(f'Model weights loaded: {memory_stats()}') - timer.record("load") - if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: - lowvram.setup_for_low_vram(sd_model, shared.cmd_opts.medvram) - else: - sd_model.to(devices.device) - timer.record("move") - shared.log.debug(f'Model weights moved: {memory_stats()}') - sd_hijack.model_hijack.hijack(sd_model) - timer.record("hijack") - sd_model.eval() - if op == 'refiner': - model_data.sd_refiner = sd_model - else: - model_data.sd_model = sd_model - sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True) # Reload embeddings after model load as they may or may not fit the model - timer.record("embeddings") - script_callbacks.model_loaded_callback(sd_model) - timer.record("callbacks") - shared.log.info(f"Model loaded in {timer.summary()}") - current_checkpoint_info = None - devices.torch_gc(force=True) - shared.log.info(f'Model load finished: {memory_stats()} cached={len(checkpoints_loaded.keys())}') - - -def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model'): - load_dict = shared.opts.sd_model_dict != model_data.sd_dict - from modules import lowvram, sd_hijack - checkpoint_info = info or select_checkpoint(op=op) # are we selecting model or dictionary - next_checkpoint_info = info or select_checkpoint(op='dict' if load_dict else 'model') if load_dict else None - if checkpoint_info is None: - unload_model_weights(op=op) - return None - orig_state = copy.deepcopy(shared.state) - shared.state = shared_state.State() - shared.state.begin('load') - if load_dict: - shared.log.debug(f'Model dict: existing={sd_model is not None} target={checkpoint_info.filename} info={info}') - else: - model_data.sd_dict = 'None' - shared.log.debug(f'Load model weights: existing={sd_model is not None} target={checkpoint_info.filename} info={info}') - if sd_model is None: - sd_model = model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner - if sd_model is None: # previous model load failed - current_checkpoint_info = None - else: - current_checkpoint_info = getattr(sd_model, 'sd_checkpoint_info', None) - if current_checkpoint_info is not None and checkpoint_info is not None and current_checkpoint_info.filename == checkpoint_info.filename: - return None - if not getattr(sd_model, 'has_accelerate', False): - if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: - lowvram.send_everything_to_cpu() - else: - sd_model.to(devices.cpu) - if (reuse_dict or shared.opts.model_reuse_dict) and not getattr(sd_model, 'has_accelerate', False): - shared.log.info('Reusing previous model dictionary') - sd_hijack.model_hijack.undo_hijack(sd_model) - else: - unload_model_weights(op=op) - sd_model = None - timer = Timer() - state_dict = get_checkpoint_state_dict(checkpoint_info, timer) - checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info) - timer.record("config") - if sd_model is None or checkpoint_config != getattr(sd_model, 'used_config', None): - sd_model = None - if shared.backend == shared.Backend.ORIGINAL: - load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op) - model_data.sd_dict = shared.opts.sd_model_dict - else: - load_diffuser(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op) - if load_dict and next_checkpoint_info is not None: - model_data.sd_dict = shared.opts.sd_model_dict - shared.opts.data["sd_model_checkpoint"] = next_checkpoint_info.title - reload_model_weights(reuse_dict=True) # ok we loaded dict now lets redo and load model on top of it - shared.state.end() - shared.state = orig_state - # data['sd_model_checkpoint'] - if op == 'model' or op == 'dict': - shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title - return model_data.sd_model - else: - shared.opts.data["sd_model_refiner"] = checkpoint_info.title - return model_data.sd_refiner - - # fallback - shared.log.info(f"Loading using fallback: {op} model={checkpoint_info.title}") - try: - load_model_weights(sd_model, checkpoint_info, state_dict, timer) - except Exception: - shared.log.error("Load model failed: restoring previous") - load_model_weights(sd_model, current_checkpoint_info, None, timer) - finally: - sd_hijack.model_hijack.hijack(sd_model) - timer.record("hijack") - script_callbacks.model_loaded_callback(sd_model) - timer.record("callbacks") - if sd_model is not None and not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram and not getattr(sd_model, 'has_accelerate', False): - sd_model.to(devices.device) - timer.record("device") - shared.state.end() - shared.state = orig_state - shared.log.info(f"Load: {op} time={timer.summary()}") - return sd_model - - -def convert_to_faketensors(tensor): - fake_module = torch._subclasses.fake_tensor.FakeTensorMode(allow_non_fake_inputs=True) # pylint: disable=protected-access - if hasattr(tensor, "weight"): - tensor.weight = torch.nn.Parameter(fake_module.from_tensor(tensor.weight)) - return tensor - - -def disable_offload(sd_model): - from accelerate.hooks import remove_hook_from_module - if not getattr(sd_model, 'has_accelerate', False): - return - for _name, model in sd_model.components.items(): - if not isinstance(model, torch.nn.Module): - continue - remove_hook_from_module(model, recurse=True) - - -def unload_model_weights(op='model', change_from='none'): - if shared.compiled_model_state is not None: - shared.compiled_model_state.compiled_cache.clear() - shared.compiled_model_state.partitioned_modules.clear() - if op == 'model' or op == 'dict': - if model_data.sd_model: - if (shared.backend == shared.Backend.ORIGINAL and change_from != shared.Backend.DIFFUSERS) or change_from == shared.Backend.ORIGINAL: - from modules import sd_hijack - model_data.sd_model.to(devices.cpu) - sd_hijack.model_hijack.undo_hijack(model_data.sd_model) - elif not (shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): - disable_offload(model_data.sd_model) - model_data.sd_model.to('meta') - model_data.sd_model = None - shared.log.debug(f'Unload weights {op}: {memory_stats()}') - else: - if model_data.sd_refiner: - if (shared.backend == shared.Backend.ORIGINAL and change_from != shared.Backend.DIFFUSERS) or change_from == shared.Backend.ORIGINAL: - from modules import sd_hijack - model_data.sd_model.to(devices.cpu) - sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner) - else: - disable_offload(model_data.sd_model) - model_data.sd_refiner.to('meta') - model_data.sd_refiner = None - shared.log.debug(f'Unload weights {op}: {memory_stats()}') - devices.torch_gc(force=True) - - -def apply_token_merging(sd_model, token_merging_ratio=0): - current_token_merging_ratio = getattr(sd_model, 'applied_token_merged_ratio', 0) - if token_merging_ratio is None or current_token_merging_ratio is None or current_token_merging_ratio == token_merging_ratio: - return - try: - if current_token_merging_ratio > 0: - tomesd.remove_patch(sd_model) - except Exception: - pass - if token_merging_ratio > 0: - if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: - shared.log.warning('Token merging not supported with HyperTile for UNet') - return - try: - tomesd.apply_patch( - sd_model, - ratio=token_merging_ratio, - use_rand=False, # can cause issues with some samplers - merge_attn=True, - merge_crossattn=False, - merge_mlp=False - ) - shared.log.info(f'Applying token merging: ratio={token_merging_ratio}') - sd_model.applied_token_merged_ratio = token_merging_ratio - except Exception: - shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') - else: - sd_model.applied_token_merged_ratio = 0 +import re +import io +import sys +import json +import time +import copy +import logging +import contextlib +import collections +import os.path +from os import mkdir +from urllib import request +from enum import Enum +from rich import progress # pylint: disable=redefined-builtin +import torch +import safetensors.torch +import diffusers +from omegaconf import OmegaConf +import tomesd +from transformers import logging as transformers_logging +import ldm.modules.midas as midas +from ldm.util import instantiate_from_config +from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_inpainting +from modules.timer import Timer +from modules.memstats import memory_stats +from modules.paths import models_path, script_path +from modules.modeldata import model_data + + +transformers_logging.set_verbosity_error() +model_dir = "Stable-diffusion" +model_path = os.path.abspath(os.path.join(paths.models_path, model_dir)) +checkpoints_list = {} +checkpoint_aliases = {} +checkpoints_loaded = collections.OrderedDict() +sd_metadata_file = os.path.join(paths.data_path, "metadata.json") +sd_metadata = None +sd_metadata_pending = 0 +sd_metadata_timer = 0 + + +class CheckpointInfo: + def __init__(self, filename): + self.name = None + self.hash = None + self.filename = filename + self.type = '' + relname = filename + app_path = os.path.abspath(script_path) + + def rel(fn, path): + try: + return os.path.relpath(fn, path) + except Exception: + return fn + + if relname.startswith('..'): + relname = os.path.abspath(relname) + if relname.startswith(shared.opts.ckpt_dir): + relname = rel(filename, shared.opts.ckpt_dir) + elif relname.startswith(shared.opts.diffusers_dir): + relname = rel(filename, shared.opts.diffusers_dir) + elif relname.startswith(model_path): + relname = rel(filename, model_path) + elif relname.startswith(script_path): + relname = rel(filename, script_path) + elif relname.startswith(app_path): + relname = rel(filename, app_path) + else: + relname = os.path.abspath(relname) + relname, ext = os.path.splitext(relname) + ext = ext.lower()[1:] + + if os.path.isfile(filename): # ckpt or safetensor + self.name = relname + self.filename = filename + self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{relname}") + self.type = ext + # self.model_name = os.path.splitext(name.replace("/", "_").replace("\\", "_"))[0] + else: # maybe a diffuser + repo = [r for r in modelloader.diffuser_repos if filename == r['name']] + if len(repo) == 0: + self.name = relname + self.filename = filename + self.sha256 = None + self.type = 'unknown' + else: + self.name = os.path.join(os.path.basename(shared.opts.diffusers_dir), repo[0]['name']) + self.filename = repo[0]['path'] + self.sha256 = repo[0]['hash'] + self.type = 'diffusers' + + self.shorthash = self.sha256[0:10] if self.sha256 else None + self.title = self.name if self.shorthash is None else f'{self.name} [{self.shorthash}]' + self.path = self.filename + self.model_name = os.path.basename(self.name) + self.metadata = read_metadata_from_safetensors(filename) + # shared.log.debug(f'Checkpoint: type={self.type} name={self.name} filename={self.filename} hash={self.shorthash} title={self.title}') + + def register(self): + checkpoints_list[self.title] = self + for i in [self.name, self.filename, self.shorthash, self.title]: + if i is not None: + checkpoint_aliases[i] = self + + def calculate_shorthash(self): + self.sha256 = hashes.sha256(self.filename, f"checkpoint/{self.name}") + if self.sha256 is None: + return None + self.shorthash = self.sha256[0:10] + checkpoints_list.pop(self.title) + self.title = f'{self.name} [{self.shorthash}]' + self.register() + return self.shorthash + + +class NoWatermark: + def apply_watermark(self, img): + return img + + +def setup_model(): + if not os.path.exists(model_path): + os.makedirs(model_path, exist_ok=True) + list_models() + enable_midas_autodownload() + + +def checkpoint_tiles(use_short=False): # pylint: disable=unused-argument + def convert(name): + return int(name) if name.isdigit() else name.lower() + def alphanumeric_key(key): + return [convert(c) for c in re.split('([0-9]+)', key)] + return sorted([x.title for x in checkpoints_list.values()], key=alphanumeric_key) + + +def list_models(): + t0 = time.time() + global checkpoints_list # pylint: disable=global-statement + checkpoints_list.clear() + checkpoint_aliases.clear() + if shared.opts.sd_disable_ckpt or shared.backend == shared.Backend.DIFFUSERS: + ext_filter = [".safetensors"] + else: + ext_filter = [".ckpt", ".safetensors"] + model_list = list(modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) + if shared.backend == shared.Backend.DIFFUSERS: + model_list += modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir, clear=True) + for filename in sorted(model_list, key=str.lower): + checkpoint_info = CheckpointInfo(filename) + if checkpoint_info.name is not None: + checkpoint_info.register() + if shared.cmd_opts.ckpt is not None: + if not os.path.exists(shared.cmd_opts.ckpt) and shared.backend == shared.Backend.ORIGINAL: + if shared.cmd_opts.ckpt.lower() != "none": + shared.log.warning(f"Requested checkpoint not found: {shared.cmd_opts.ckpt}") + else: + checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt) + if checkpoint_info.name is not None: + checkpoint_info.register() + shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title + elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None: + shared.log.warning(f"Checkpoint not found: {shared.cmd_opts.ckpt}") + shared.log.info(f'Available models: path="{shared.opts.ckpt_dir}" items={len(checkpoints_list)} time={time.time()-t0:.2f}') + + checkpoints_list = dict(sorted(checkpoints_list.items(), key=lambda cp: cp[1].filename)) + """ + if len(checkpoints_list) == 0: + if not shared.cmd_opts.no_download: + key = input('Download the default model? (y/N) ') + if key.lower().startswith('y'): + if shared.backend == shared.Backend.ORIGINAL: + model_url = "https://huggingface.co/runwayml/stable-diffusion-v1-5/resolve/main/v1-5-pruned-emaonly.safetensors" + shared.opts.data['sd_model_checkpoint'] = "v1-5-pruned-emaonly.safetensors" + model_list = modelloader.load_models(model_path=model_path, model_url=model_url, command_path=shared.opts.ckpt_dir, ext_filter=[".ckpt", ".safetensors"], download_name="v1-5-pruned-emaonly.safetensors", ext_blacklist=[".vae.ckpt", ".vae.safetensors"]) + else: + default_model_id = "runwayml/stable-diffusion-v1-5" + modelloader.download_diffusers_model(default_model_id, shared.opts.diffusers_dir) + model_list = modelloader.load_diffusers_models(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir) + + for filename in sorted(model_list, key=str.lower): + checkpoint_info = CheckpointInfo(filename) + if checkpoint_info.name is not None: + checkpoint_info.register() + """ + +def update_model_hashes(): + txt = [] + lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.hash is None] + # shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') + for ckpt in lst: + ckpt.hash = model_hash(ckpt.filename) + # txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') + # txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') + lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.sha256 is None or ckpt.shorthash is None] + shared.log.info(f'Models list: hash missing={len(lst)} total={len(checkpoints_list)}') + for ckpt in lst: + ckpt.sha256 = hashes.sha256(ckpt.filename, f"checkpoint/{ckpt.name}") + ckpt.shorthash = ckpt.sha256[0:10] if ckpt.sha256 is not None else None + if ckpt.sha256 is not None: + txt.append(f'Calculated full hash: {ckpt.title} {ckpt.shorthash}') + else: + txt.append(f'Skipped hash calculation: {ckpt.title}') + txt.append(f'Updated hashes for {len(lst)} out of {len(checkpoints_list)} models') + txt = '
'.join(txt) + return txt + + +def get_closet_checkpoint_match(search_string): + checkpoint_info = checkpoint_aliases.get(search_string, None) + if checkpoint_info is not None: + return checkpoint_info + found = sorted([info for info in checkpoints_list.values() if search_string in info.title], key=lambda x: len(x.title)) + if found: + return found[0] + found = sorted([info for info in checkpoints_list.values() if search_string.split(' ')[0] in info.title], key=lambda x: len(x.title)) + if found: + return found[0] + return None + + +def model_hash(filename): + """old hash that only looks at a small part of the file and is prone to collisions""" + try: + with open(filename, "rb") as file: + import hashlib + # t0 = time.time() + m = hashlib.sha256() + file.seek(0x100000) + m.update(file.read(0x10000)) + shorthash = m.hexdigest()[0:8] + # t1 = time.time() + # shared.log.debug(f'Calculating short hash: {filename} hash={shorthash} time={(t1-t0):.2f}') + return shorthash + except FileNotFoundError: + return 'NOFILE' + except Exception: + return 'NOHASH' + + +def select_checkpoint(op='model'): + if op == 'dict': + model_checkpoint = shared.opts.sd_model_dict + elif op == 'refiner': + model_checkpoint = shared.opts.data.get('sd_model_refiner', None) + else: + model_checkpoint = shared.opts.sd_model_checkpoint + if model_checkpoint is None or model_checkpoint == 'None': + return None + checkpoint_info = get_closet_checkpoint_match(model_checkpoint) + if checkpoint_info is not None: + shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"') + return checkpoint_info + if len(checkpoints_list) == 0 and not shared.cmd_opts.no_download: + shared.log.warning("Cannot generate without a checkpoint") + shared.log.info("Set system paths to use existing folders in a different location") + shared.log.info("Or use --ckpt to force using existing checkpoint") + return None + checkpoint_info = next(iter(checkpoints_list.values())) + if model_checkpoint is not None: + if model_checkpoint != 'model.ckpt' and model_checkpoint != 'runwayml/stable-diffusion-v1-5': + shared.log.warning(f"Selected checkpoint not found: {model_checkpoint}") + else: + shared.log.info("Selecting first available checkpoint") + # shared.log.warning(f"Loading fallback checkpoint: {checkpoint_info.title}") + shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title + shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"') + return checkpoint_info + + +checkpoint_dict_replacements = { + 'cond_stage_model.transformer.embeddings.': 'cond_stage_model.transformer.text_model.embeddings.', + 'cond_stage_model.transformer.encoder.': 'cond_stage_model.transformer.text_model.encoder.', + 'cond_stage_model.transformer.final_layer_norm.': 'cond_stage_model.transformer.text_model.final_layer_norm.', +} + + +def transform_checkpoint_dict_key(k): + for text, replacement in checkpoint_dict_replacements.items(): + if k.startswith(text): + k = replacement + k[len(text):] + return k + + +def get_state_dict_from_checkpoint(pl_sd): + pl_sd = pl_sd.pop("state_dict", pl_sd) + pl_sd.pop("state_dict", None) + sd = {} + for k, v in pl_sd.items(): + new_key = transform_checkpoint_dict_key(k) + if new_key is not None: + sd[new_key] = v + pl_sd.clear() + pl_sd.update(sd) + return pl_sd + + +def write_metadata(): + global sd_metadata_pending # pylint: disable=global-statement + if sd_metadata_pending == 0: + shared.log.debug(f'Model metadata: file="{sd_metadata_file}" no changes') + return + shared.writefile(sd_metadata, sd_metadata_file) + shared.log.info(f'Model metadata saved: file="{sd_metadata_file}" items={sd_metadata_pending} time={sd_metadata_timer:.2f}') + sd_metadata_pending = 0 + + +def scrub_dict(dict_obj, keys): + for key in list(dict_obj.keys()): + if not isinstance(dict_obj, dict): + continue + if key in keys: + dict_obj.pop(key, None) + elif isinstance(dict_obj[key], dict): + scrub_dict(dict_obj[key], keys) + elif isinstance(dict_obj[key], list): + for item in dict_obj[key]: + scrub_dict(item, keys) + + +def read_metadata_from_safetensors(filename): + global sd_metadata # pylint: disable=global-statement + if sd_metadata is None: + if not os.path.isfile(sd_metadata_file): + sd_metadata = {} + else: + sd_metadata = shared.readfile(sd_metadata_file, lock=True) + res = sd_metadata.get(filename, None) + if res is not None: + return res + if not filename.endswith(".safetensors"): + return {} + if shared.cmd_opts.no_metadata: + return {} + res = {} + try: + t0 = time.time() + with open(filename, mode="rb") as file: + metadata_len = file.read(8) + metadata_len = int.from_bytes(metadata_len, "little") + json_start = file.read(2) + if metadata_len <= 2 or json_start not in (b'{"', b"{'"): + shared.log.error(f"Not a valid safetensors file: {filename}") + json_data = json_start + file.read(metadata_len-2) + json_obj = json.loads(json_data) + for k, v in json_obj.get("__metadata__", {}).items(): + if v.startswith("data:"): + v = 'data' + if k == 'format' and v == 'pt': + continue + large = True if len(v) > 2048 else False + if large and k == 'ss_datasets': + continue + if large and k == 'workflow': + continue + if large and k == 'prompt': + continue + if large and k == 'ss_bucket_info': + continue + if v[0:1] == '{': + try: + v = json.loads(v) + if large and k == 'ss_tag_frequency': + v = { i: len(j) for i, j in v.items() } + if large and k == 'sd_merge_models': + scrub_dict(v, ['sd_merge_recipe']) + except Exception: + pass + res[k] = v + sd_metadata[filename] = res + global sd_metadata_pending # pylint: disable=global-statement + sd_metadata_pending += 1 + t1 = time.time() + global sd_metadata_timer # pylint: disable=global-statement + sd_metadata_timer += (t1 - t0) + except Exception as e: + shared.log.error(f"Error reading metadata from: {filename} {e}") + return res + + +def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unused-argument + if not os.path.isfile(checkpoint_file): + shared.log.error(f"Model is not a file: {checkpoint_file}") + return None + try: + pl_sd = None + with progress.open(checkpoint_file, 'rb', description=f'[cyan]Loading model: [yellow]{checkpoint_file}', auto_refresh=True, console=shared.console) as f: + _, extension = os.path.splitext(checkpoint_file) + if extension.lower() == ".ckpt" and shared.opts.sd_disable_ckpt: + shared.log.warning(f"Checkpoint loading disabled: {checkpoint_file}") + return None + if shared.opts.stream_load: + if extension.lower() == ".safetensors": + # shared.log.debug('Model weights loading: type=safetensors mode=buffered') + buffer = f.read() + pl_sd = safetensors.torch.load(buffer) + else: + # shared.log.debug('Model weights loading: type=checkpoint mode=buffered') + buffer = io.BytesIO(f.read()) + pl_sd = torch.load(buffer, map_location='cpu') + else: + if extension.lower() == ".safetensors": + # shared.log.debug('Model weights loading: type=safetensors mode=mmap') + pl_sd = safetensors.torch.load_file(checkpoint_file, device='cpu') + else: + # shared.log.debug('Model weights loading: type=checkpoint mode=direct') + pl_sd = torch.load(f, map_location='cpu') + sd = get_state_dict_from_checkpoint(pl_sd) + del pl_sd + except Exception as e: + errors.display(e, f'Load model: {checkpoint_file}') + sd = None + return sd + + +def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer): + if not os.path.isfile(checkpoint_info.filename): + return None + if checkpoint_info in checkpoints_loaded: + shared.log.info("Model weights loading: from cache") + return checkpoints_loaded[checkpoint_info] + res = read_state_dict(checkpoint_info.filename) + timer.record("load") + return res + + +def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo, state_dict, timer): + _pipeline, _model_type = detect_pipeline(checkpoint_info.path, 'model') + shared.log.debug(f'Model weights loading: {memory_stats()}') + timer.record("hash") + if model_data.sd_dict == 'None': + shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title + if state_dict is None: + state_dict = get_checkpoint_state_dict(checkpoint_info, timer) + try: + model.load_state_dict(state_dict, strict=False) + except Exception as e: + shared.log.error(f'Error loading model weights: {checkpoint_info.filename}') + shared.log.error(' '.join(str(e).splitlines()[:2])) + return False + del state_dict + timer.record("apply") + if shared.opts.sd_checkpoint_cache > 0: + # cache newly loaded model + checkpoints_loaded[checkpoint_info] = model.state_dict().copy() + if shared.opts.opt_channelslast: + model.to(memory_format=torch.channels_last) + timer.record("channels") + if not shared.opts.no_half: + vae = model.first_stage_model + depth_model = getattr(model, 'depth_model', None) + # with --no-half-vae, remove VAE from model when doing half() to prevent its weights from being converted to float16 + if shared.opts.no_half_vae: + model.first_stage_model = None + # with --upcast-sampling, don't convert the depth model weights to float16 + if shared.opts.upcast_sampling and depth_model: + model.depth_model = None + model.half() + model.first_stage_model = vae + if depth_model: + model.depth_model = depth_model + if shared.opts.cuda_cast_unet: + devices.dtype_unet = model.model.diffusion_model.dtype + else: + model.model.diffusion_model.to(devices.dtype_unet) + model.first_stage_model.to(devices.dtype_vae) + # clean up cache if limit is reached + while len(checkpoints_loaded) > shared.opts.sd_checkpoint_cache: + checkpoints_loaded.popitem(last=False) + model.sd_model_hash = checkpoint_info.calculate_shorthash() + model.sd_model_checkpoint = checkpoint_info.filename + model.sd_checkpoint_info = checkpoint_info + model.is_sdxl = False # a1111 compatibility item + model.is_sd2 = hasattr(model.cond_stage_model, 'model') # a1111 compatibility item + model.is_sd1 = not hasattr(model.cond_stage_model, 'model') # a1111 compatibility item + model.logvar = model.logvar.to(devices.device) if hasattr(model, 'logvar') else None # fix for training + shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256 + sd_vae.delete_base_vae() + sd_vae.clear_loaded_vae() + vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) + sd_vae.load_vae(model, vae_file, vae_source) + timer.record("vae") + return True + + +def enable_midas_autodownload(): + """ + Gives the ldm.modules.midas.api.load_model function automatic downloading. + + When the 512-depth-ema model, and other future models like it, is loaded, + it calls midas.api.load_model to load the associated midas depth model. + This function applies a wrapper to download the model to the correct + location automatically. + """ + midas_path = os.path.join(paths.models_path, 'midas') + for k, v in midas.api.ISL_PATHS.items(): + file_name = os.path.basename(v) + midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) + midas_urls = { + "dpt_large": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_large-midas-2f21e586.pt", + "dpt_hybrid": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_hybrid-midas-501f0c75.pt", + "midas_v21": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21-f6b98070.pt", + "midas_v21_small": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21_small-70d6b9c8.pt", + } + midas.api.load_model_inner = midas.api.load_model + + def load_model_wrapper(model_type): + path = midas.api.ISL_PATHS[model_type] + if not os.path.exists(path): + if not os.path.exists(midas_path): + mkdir(midas_path) + shared.log.info(f"Downloading midas model weights for {model_type} to {path}") + request.urlretrieve(midas_urls[model_type], path) + shared.log.info(f"{model_type} downloaded") + return midas.api.load_model_inner(model_type) + + midas.api.load_model = load_model_wrapper + + +def repair_config(sd_config): + if "use_ema" not in sd_config.model.params: + sd_config.model.params.use_ema = False + if shared.opts.no_half: + sd_config.model.params.unet_config.params.use_fp16 = False + elif shared.opts.upcast_sampling: + sd_config.model.params.unet_config.params.use_fp16 = True if sys.platform != 'darwin' else False + if getattr(sd_config.model.params.first_stage_config.params.ddconfig, "attn_type", None) == "vanilla-xformers" and not shared.xformers_available: + sd_config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla" + # For UnCLIP-L, override the hardcoded karlo directory + if "noise_aug_config" in sd_config.model.params and "clip_stats_path" in sd_config.model.params.noise_aug_config.params: + karlo_path = os.path.join(paths.models_path, 'karlo') + sd_config.model.params.noise_aug_config.params.clip_stats_path = sd_config.model.params.noise_aug_config.params.clip_stats_path.replace("checkpoints/karlo_models", karlo_path) + + +sd1_clip_weight = 'cond_stage_model.transformer.text_model.embeddings.token_embedding.weight' +sd2_clip_weight = 'cond_stage_model.model.transformer.resblocks.0.attn.in_proj_weight' + + +def change_backend(): + shared.log.info(f'Backend changed: {shared.backend}') + shared.log.warning('Full server restart required to apply all changes') + if shared.backend == shared.Backend.ORIGINAL: + change_from = shared.Backend.DIFFUSERS + else: + change_from = shared.Backend.ORIGINAL + unload_model_weights(change_from=change_from) + checkpoints_loaded.clear() + from modules.sd_samplers import list_samplers + list_samplers(shared.backend) + list_models() + from modules.sd_vae import refresh_vae_list + refresh_vae_list() + + +def detect_pipeline(f: str, op: str = 'model', warning=True): + if not f.endswith('.safetensors'): + return None, None + guess = shared.opts.diffusers_pipeline + warn = shared.log.warning if warning else lambda *args, **kwargs: None + if guess == 'Autodetect': + try: + # guess by size + size = round(os.path.getsize(f) / 1024 / 1024) + if size < 128: + warn(f'Model size smaller than expected: {f} size={size} MB') + elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160 + warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB') + guess = 'VAE' + elif size >= 5351 and size <= 5359: # 5353 + guess = 'Stable Diffusion' # SD v2 + elif size >= 5791 and size <= 5799: # 5795 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as SD-XL refiner model, but attempting to load using backend=original: {op}={f} size={size} MB') + if op == 'model': + warn(f'Model detected as SD-XL refiner model, but attempting to load a base model: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL' + elif (size >= 6611 and size <= 6619) or (size >= 6771 and size <= 6779): # 6617, HassakuXL is 6776 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as SD-XL base model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL' + elif size >= 3361 and size <= 3369: # 3368 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as SD upscale model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Stable Diffusion Upscale' + elif size >= 4891 and size <= 4899: # 4897 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as SD XL inpaint model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL Inpaint' + elif size >= 9791 and size <= 9799: # 9794 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as SD XL instruct pix2pix model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL Instruct' + elif size > 3138 and size < 3142: #3140 + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as Segmind Vega model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL' + else: + guess = 'Stable Diffusion' + # guess by name + """ + if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper(): + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as LCM model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Latent Consistency Model' + """ + if 'PixArt' in f: + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as PixArt Alpha model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'PixArt Alpha' + # switch for specific variant + if guess == 'Stable Diffusion' and 'inpaint' in f.lower(): + guess = 'Stable Diffusion Inpaint' + elif guess == 'Stable Diffusion' and 'instruct' in f.lower(): + guess = 'Stable Diffusion Instruct' + if guess == 'Stable Diffusion XL' and 'inpaint' in f.lower(): + guess = 'Stable Diffusion XL Inpaint' + elif guess == 'Stable Diffusion XL' and 'instruct' in f.lower(): + guess = 'Stable Diffusion XL Instruct' + # get actual pipeline + pipeline = shared_items.get_pipelines().get(guess, None) + shared.log.info(f'Autodetect: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB') + except Exception as e: + shared.log.error(f'Error detecting diffusers pipeline: model={f} {e}') + return None, None + else: + try: + size = round(os.path.getsize(f) / 1024 / 1024) + pipeline = shared_items.get_pipelines().get(guess, None) + shared.log.info(f'Diffusers: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB') + except Exception as e: + shared.log.error(f'Error loading diffusers pipeline: model={f} {e}') + + if pipeline is None: + shared.log.warning(f'Autodetect: pipeline not recognized: {guess}: {op}={f} size={size}') + pipeline = diffusers.StableDiffusionPipeline + return pipeline, guess + + +def copy_diffuser_options(new_pipe, orig_pipe): + new_pipe.sd_checkpoint_info = orig_pipe.sd_checkpoint_info + new_pipe.sd_model_checkpoint = orig_pipe.sd_model_checkpoint + new_pipe.embedding_db = getattr(orig_pipe, 'embedding_db', None) + new_pipe.sd_model_hash = getattr(orig_pipe, 'sd_model_hash', None) + new_pipe.has_accelerate = getattr(orig_pipe, 'has_accelerate', False) + new_pipe.is_sdxl = getattr(orig_pipe, 'is_sdxl', False) # a1111 compatibility item + new_pipe.is_sd2 = getattr(orig_pipe, 'is_sd2', False) + new_pipe.is_sd1 = getattr(orig_pipe, 'is_sd1', True) + + +def set_diffuser_options(sd_model, vae = None, op: str = 'model'): + if sd_model is None: + shared.log.warning(f'{op} is not loaded') + return + if (shared.opts.diffusers_model_cpu_offload or shared.cmd_opts.medvram) and (shared.opts.diffusers_seq_cpu_offload or shared.cmd_opts.lowvram): + shared.log.warning(f'Setting {op}: Model CPU offload and Sequential CPU offload are not compatible') + shared.log.debug(f'Setting {op}: disabling model CPU offload') + shared.opts.diffusers_model_cpu_offload=False + shared.cmd_opts.medvram=False + + if hasattr(sd_model, "watermark"): + sd_model.watermark = NoWatermark() + sd_model.has_accelerate = False + if hasattr(sd_model, "vae"): + if vae is not None: + sd_model.vae = vae + if shared.opts.diffusers_vae_upcast != 'default': + if shared.opts.diffusers_vae_upcast == 'true': + sd_model.vae.config.force_upcast = True + else: + sd_model.vae.config.force_upcast = False + if shared.opts.no_half_vae: + devices.dtype_vae = torch.float32 + sd_model.vae.to(devices.dtype_vae) + shared.log.debug(f'Setting {op} VAE: name={sd_vae.loaded_vae_file} upcast={sd_model.vae.config.get("force_upcast", None)}') + if hasattr(sd_model, "enable_model_cpu_offload"): + if (shared.cmd_opts.medvram and devices.backend != "directml") or shared.opts.diffusers_model_cpu_offload: + shared.log.debug(f'Setting {op}: enable model CPU offload') + if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: + shared.opts.diffusers_move_base = False + shared.opts.diffusers_move_unet = False + shared.opts.diffusers_move_refiner = False + shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled') + sd_model.enable_model_cpu_offload() + sd_model.has_accelerate = True + if hasattr(sd_model, "enable_sequential_cpu_offload"): + if shared.cmd_opts.lowvram or shared.opts.diffusers_seq_cpu_offload: + shared.log.debug(f'Setting {op}: enable sequential CPU offload') + if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: + shared.opts.diffusers_move_base = False + shared.opts.diffusers_move_unet = False + shared.opts.diffusers_move_refiner = False + shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled') + sd_model.enable_sequential_cpu_offload(device=devices.device) + sd_model.has_accelerate = True + if hasattr(sd_model, "enable_vae_slicing"): + if shared.cmd_opts.lowvram or shared.opts.diffusers_vae_slicing: + shared.log.debug(f'Setting {op}: enable VAE slicing') + sd_model.enable_vae_slicing() + else: + sd_model.disable_vae_slicing() + if hasattr(sd_model, "enable_vae_tiling"): + if shared.cmd_opts.lowvram or shared.opts.diffusers_vae_tiling: + shared.log.debug(f'Setting {op}: enable VAE tiling') + sd_model.enable_vae_tiling() + else: + sd_model.disable_vae_tiling() + if hasattr(sd_model, "enable_attention_slicing"): + if shared.cmd_opts.lowvram or shared.opts.diffusers_attention_slicing: + shared.log.debug(f'Setting {op}: enable attention slicing') + sd_model.enable_attention_slicing() + else: + sd_model.disable_attention_slicing() + if hasattr(sd_model, "vqvae"): + sd_model.vqvae.to(torch.float32) # vqvae is producing nans in fp16 + if shared.opts.cross_attention_optimization == "xFormers" and hasattr(sd_model, 'enable_xformers_memory_efficient_attention'): + sd_model.enable_xformers_memory_efficient_attention() + if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'fuse_qkv_projections'): + shared.log.debug(f'Setting {op}: enable fused projections') + sd_model.fuse_qkv_projections() + if shared.opts.diffusers_eval: + if hasattr(sd_model, "unet") and hasattr(sd_model.unet, "requires_grad_"): + sd_model.unet.requires_grad_(False) + sd_model.unet.eval() + if hasattr(sd_model, "vae") and hasattr(sd_model.vae, "requires_grad_"): + sd_model.vae.requires_grad_(False) + sd_model.vae.eval() + if hasattr(sd_model, "text_encoder") and hasattr(sd_model.text_encoder, "requires_grad_"): + sd_model.text_encoder.requires_grad_(False) + sd_model.text_encoder.eval() + if shared.opts.diffusers_quantization: + sd_model = sd_models_compile.dynamic_quantization(sd_model) + + if shared.opts.opt_channelslast and hasattr(sd_model, 'unet'): + shared.log.debug(f'Setting {op}: enable channels last') + sd_model.unet.to(memory_format=torch.channels_last) + + +def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): # pylint: disable=unused-argument + import torch # pylint: disable=reimported,redefined-outer-name + if shared.cmd_opts.profile: + import cProfile + pr = cProfile.Profile() + pr.enable() + if timer is None: + timer = Timer() + logging.getLogger("diffusers").setLevel(logging.ERROR) + timer.record("diffusers") + devices.set_cuda_params() + diffusers_load_config = { + "low_cpu_mem_usage": True, + "torch_dtype": devices.dtype, + "safety_checker": None, + "requires_safety_checker": False, + "load_safety_checker": False, + "load_connected_pipeline": True, + # TODO: use_safetensors cant enable for all checkpoints just yet + } + if shared.opts.diffusers_model_load_variant == 'default': + if devices.dtype == torch.float16: + diffusers_load_config['variant'] = 'fp16' + elif shared.opts.diffusers_model_load_variant == 'fp32': + pass + else: + diffusers_load_config['variant'] = shared.opts.diffusers_model_load_variant + + if shared.opts.diffusers_pipeline == 'Custom Diffusers Pipeline' and len(shared.opts.custom_diffusers_pipeline) > 0: + shared.log.debug(f'Diffusers custom pipeline: {shared.opts.custom_diffusers_pipeline}') + diffusers_load_config['custom_pipeline'] = shared.opts.custom_diffusers_pipeline + + # if 'LCM' in checkpoint_info.path: + # diffusers_load_config['custom_pipeline'] = 'latent_consistency_txt2img' + + if shared.opts.data.get('sd_model_checkpoint', '') == 'model.ckpt' or shared.opts.data.get('sd_model_checkpoint', '') == '': + shared.opts.data['sd_model_checkpoint'] = "runwayml/stable-diffusion-v1-5" + + if op == 'model' or op == 'dict': + if (model_data.sd_model is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model + return + else: + if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model + return + + sd_model = None + + try: + if shared.cmd_opts.ckpt is not None and os.path.isdir(shared.cmd_opts.ckpt) and model_data.initial: # initial load + ckpt_basename = os.path.basename(shared.cmd_opts.ckpt) + model_name = modelloader.find_diffuser(ckpt_basename) + if model_name is not None: + shared.log.info(f'Load model {op}: {model_name}') + model_file = modelloader.download_diffusers_model(hub_id=model_name) + try: + shared.log.debug(f'Model load {op} config: {diffusers_load_config}') + sd_model = diffusers.DiffusionPipeline.from_pretrained(model_file, **diffusers_load_config) + except Exception as e: + shared.log.error(f'Failed loading model: {model_file} {e}') + list_models() # rescan for downloaded model + checkpoint_info = CheckpointInfo(model_name) + + checkpoint_info = checkpoint_info or select_checkpoint(op=op) + if checkpoint_info is None: + unload_model_weights(op=op) + return + + vae = None + sd_vae.loaded_vae_file = None + if op == 'model' or op == 'refiner': + vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) + vae = sd_vae.load_vae_diffusers(checkpoint_info.path, vae_file, vae_source) + if vae is not None: + diffusers_load_config["vae"] = vae + + shared.log.debug(f'Diffusers loading: path="{checkpoint_info.path}"') + if os.path.isdir(checkpoint_info.path): + err1 = None + err2 = None + err3 = None + try: # try autopipeline first, best choice but not all pipelines are available + sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except Exception as e: + err1 = e + # shared.log.error(f'AutoPipeline: {e}') + try: # try diffusion pipeline next second-best choice, works for most non-linked pipelines + if err1 is not None: + sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except Exception as e: + err2 = e + # shared.log.error(f'DiffusionPipeline: {e}') + try: # try basic pipeline next just in case + if err2 is not None: + sd_model = diffusers.StableDiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except Exception as e: + err3 = e # ignore last error + shared.log.error(f'StableDiffusionPipeline: {e}') + if err3 is not None: + shared.log.error(f'Failed loading {op}: {checkpoint_info.path} auto={err1} diffusion={err2}') + return + elif os.path.isfile(checkpoint_info.path) and checkpoint_info.path.lower().endswith('.safetensors'): + # diffusers_load_config["local_files_only"] = True + diffusers_load_config["extract_ema"] = shared.opts.diffusers_extract_ema + pipeline, model_type = detect_pipeline(checkpoint_info.path, op) + if pipeline is None: + shared.log.error(f'Diffusers {op} pipeline not initialized: {shared.opts.diffusers_pipeline}') + return + try: + if model_type.startswith('Stable Diffusion'): + diffusers_load_config['force_zeros_for_empty_prompt '] = shared.opts.diffusers_force_zeros + diffusers_load_config['requires_aesthetics_score'] = shared.opts.diffusers_aesthetics_score + if 'inpainting' in checkpoint_info.path.lower(): + diffusers_load_config['config_files'] = { + 'v1': 'configs/v1-inpainting-inference.yaml', + 'v2': 'configs/v2-inference-768-v.yaml', + 'xl': 'configs/sd_xl_base.yaml', + 'xl_refiner': 'configs/sd_xl_refiner.yaml', + } + else: + diffusers_load_config['config_files'] = { + 'v1': 'configs/v1-inference.yaml', + 'v2': 'configs/v2-inference-768-v.yaml', + 'xl': 'configs/sd_xl_base.yaml', + 'xl_refiner': 'configs/sd_xl_refiner.yaml', + } + if hasattr(pipeline, 'from_single_file'): + diffusers_load_config['use_safetensors'] = True + sd_model = pipeline.from_single_file(checkpoint_info.path, **diffusers_load_config) + if sd_model is not None and hasattr(sd_model, 'unet') and hasattr(sd_model.unet, 'config') and 'inpainting' in checkpoint_info.path.lower(): + shared.log.debug('Model patch: type=inpaint') + sd_model.unet.config.in_channels = 9 + elif hasattr(pipeline, 'from_ckpt'): + sd_model = pipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config) + else: + shared.log.error(f'Diffusers {op} cannot load safetensor model: {checkpoint_info.path} {shared.opts.diffusers_pipeline}') + return + if sd_model is not None: + diffusers_load_config.pop('vae', None) + diffusers_load_config.pop('safety_checker', None) + diffusers_load_config.pop('requires_safety_checker', None) + diffusers_load_config.pop('load_safety_checker', None) + diffusers_load_config.pop('config_files', None) + diffusers_load_config.pop('local_files_only', None) + shared.log.debug(f'Setting {op}: pipeline={sd_model.__class__.__name__} config={diffusers_load_config}') # pylint: disable=protected-access + except Exception as e: + shared.log.error(f'Diffusers failed loading: {op}={checkpoint_info.path} pipeline={shared.opts.diffusers_pipeline}/{sd_model.__class__.__name__} {e}') + errors.display(e, f'loading {op}={checkpoint_info.path} pipeline={shared.opts.diffusers_pipeline}/{sd_model.__class__.__name__}') + return + else: + shared.log.error(f'Diffusers cannot load: {op}={checkpoint_info.path}') + return + + if "StableDiffusion" in sd_model.__class__.__name__: + pass # scheduler is created on first use + elif "Kandinsky" in sd_model.__class__.__name__: + sd_model.scheduler.name = 'DDIM' + + set_diffuser_options(sd_model, vae, op) + + base_sent_to_cpu=False + if (shared.opts.cuda_compile and shared.opts.cuda_compile_backend != 'none') or shared.opts.ipex_optimize: + if op == 'refiner' and not getattr(sd_model, 'has_accelerate', False): + gpu_vram = memory_stats().get('gpu', {}) + free_vram = gpu_vram.get('total', 0) - gpu_vram.get('used', 0) + refiner_enough_vram = free_vram >= 7 if "StableDiffusionXL" in sd_model.__class__.__name__ else 3 + if not shared.opts.diffusers_move_base and refiner_enough_vram: + sd_model.to(devices.device) + base_sent_to_cpu=False + else: + if not refiner_enough_vram and not (shared.opts.diffusers_move_base and shared.opts.diffusers_move_refiner): + shared.log.warning(f"Insufficient GPU memory, using system memory as fallback: free={free_vram} GB") + if not shared.opts.shared.opts.diffusers_seq_cpu_offload and not shared.opts.diffusers_model_cpu_offload: + shared.log.debug('Enabled moving base model to CPU') + shared.log.debug('Enabled moving refiner model to CPU') + shared.opts.diffusers_move_base=True + shared.opts.diffusers_move_refiner=True + shared.log.debug('Moving base model to CPU') + if model_data.sd_model is not None: + model_data.sd_model.to(devices.cpu) + devices.torch_gc(force=True) + sd_model.to(devices.device) + base_sent_to_cpu=True + elif not getattr(sd_model, 'has_accelerate', False): + sd_model.to(devices.device) + + sd_models_compile.compile_diffusers(sd_model) + + if sd_model is None: + shared.log.error('Diffuser model not loaded') + return + sd_model.sd_model_hash = checkpoint_info.calculate_shorthash() # pylint: disable=attribute-defined-outside-init + sd_model.sd_checkpoint_info = checkpoint_info # pylint: disable=attribute-defined-outside-init + sd_model.sd_model_checkpoint = checkpoint_info.filename # pylint: disable=attribute-defined-outside-init + sd_model.is_sdxl = False # a1111 compatibility item + sd_model.is_sd2 = hasattr(sd_model, 'cond_stage_model') and hasattr(sd_model.cond_stage_model, 'model') # a1111 compatibility item + sd_model.is_sd1 = not sd_model.is_sd2 # a1111 compatibility item + sd_model.logvar = sd_model.logvar.to(devices.device) if hasattr(sd_model, 'logvar') else None # fix for training + shared.opts.data["sd_checkpoint_hash"] = checkpoint_info.sha256 + if hasattr(sd_model, "set_progress_bar_config"): + sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining}', ncols=80, colour='#327fba') + if op == 'refiner' and shared.opts.diffusers_move_refiner and not getattr(sd_model, 'has_accelerate', False): + shared.log.debug('Moving refiner model to CPU') + sd_model.to(devices.cpu) + elif not getattr(sd_model, 'has_accelerate', False): # In offload modes, accelerate will move models around + sd_model.to(devices.device) + if op == 'refiner' and base_sent_to_cpu: + shared.log.debug('Moving base model back to GPU') + model_data.sd_model.to(devices.device) + except Exception as e: + shared.log.error("Failed to load diffusers model") + errors.display(e, "loading Diffusers model") + + if sd_model is not None: + from modules.textual_inversion import textual_inversion + sd_model.embedding_db = textual_inversion.EmbeddingDatabase() + if op == 'refiner': + model_data.sd_refiner = sd_model + else: + model_data.sd_model = sd_model + sd_model.embedding_db.add_embedding_dir(shared.opts.embeddings_dir) + sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) + + timer.record("load") + devices.torch_gc(force=True) + if shared.cmd_opts.profile: + errors.profile(pr, 'Load') + script_callbacks.model_loaded_callback(sd_model) + shared.log.info(f"Load {op}: time={timer.summary()} native={get_native(sd_model)} {memory_stats()}") + + +class DiffusersTaskType(Enum): + TEXT_2_IMAGE = 1 + IMAGE_2_IMAGE = 2 + INPAINTING = 3 + INSTRUCT = 4 + + +def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType: + if pipe.__class__.__name__ == "StableDiffusionXLInstructPix2PixPipeline": + return DiffusersTaskType.INSTRUCT + elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING.values(): + return DiffusersTaskType.IMAGE_2_IMAGE + elif pipe.__class__ in diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values(): + return DiffusersTaskType.INPAINTING + else: + return DiffusersTaskType.TEXT_2_IMAGE + + +def switch_diffuser_pipe(pipeline, cls): + try: + new_pipe = None + if isinstance(pipeline, cls): + return pipeline + elif isinstance(pipeline, diffusers.StableDiffusionXLPipeline): + new_pipe = cls( + vae=pipeline.vae, + text_encoder=pipeline.text_encoder, + text_encoder_2=pipeline.text_encoder_2, + tokenizer=pipeline.tokenizer, + tokenizer_2=pipeline.tokenizer_2, + unet=pipeline.unet, + scheduler=pipeline.scheduler, + feature_extractor=getattr(pipeline, 'feature_extractor', None), + ).to(pipeline.device) + elif isinstance(pipeline, diffusers.StableDiffusionPipeline): + new_pipe = cls( + vae=pipeline.vae, + text_encoder=pipeline.text_encoder, + tokenizer=pipeline.tokenizer, + unet=pipeline.unet, + scheduler=pipeline.scheduler, + feature_extractor=getattr(pipeline, 'feature_extractor', None), + requires_safety_checker=False, + safety_checker=None, + ).to(pipeline.device) + else: + shared.log.error(f'Pipeline switch error: {pipeline.__class__.__name__} unrecognized') + return pipeline + if new_pipe is not None: + copy_diffuser_options(new_pipe, pipeline) + shared.log.debug(f'Pipeline switch: from={pipeline.__class__.__name__} to={new_pipe.__class__.__name__}') + return new_pipe + else: + shared.log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls.__name__} empty pipeline') + except Exception as e: + shared.log.error(f'Pipeline switch error: from={pipeline.__class__.__name__} to={cls.__name__} {e}') + return pipeline + + +def set_diffuser_pipe(pipe, new_pipe_type): + sd_checkpoint_info = getattr(pipe, "sd_checkpoint_info", None) + sd_model_checkpoint = getattr(pipe, "sd_model_checkpoint", None) + sd_model_hash = getattr(pipe, "sd_model_hash", None) + has_accelerate = getattr(pipe, "has_accelerate", None) + embedding_db = getattr(pipe, "embedding_db", None) + image_encoder = getattr(pipe, "image_encoder", None) + feature_extractor = getattr(pipe, "feature_extractor", None) + + # skip specific pipelines + if pipe.__class__.__name__ == 'StableDiffusionReferencePipeline' or pipe.__class__.__name__ == 'StableDiffusionAdapterPipeline': + return pipe + + try: + if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: + new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: + new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.INPAINTING: + new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) + except Exception as e: # pylint: disable=unused-variable + shared.log.warning(f'Failed to change: type={new_pipe_type} pipeline={pipe.__class__.__name__} {e}') + return pipe + + if pipe.__class__ == new_pipe.__class__: + return pipe + new_pipe.sd_checkpoint_info = sd_checkpoint_info + new_pipe.sd_model_checkpoint = sd_model_checkpoint + new_pipe.sd_model_hash = sd_model_hash + new_pipe.has_accelerate = has_accelerate + new_pipe.embedding_db = embedding_db + new_pipe.image_encoder = image_encoder + new_pipe.feature_extractor = feature_extractor + new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item + new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False) + new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) + shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__}") + pipe = new_pipe + return pipe + + +def get_native(pipe: diffusers.DiffusionPipeline): + if hasattr(pipe, "vae") and hasattr(pipe.vae.config, "sample_size"): + # Stable Diffusion + size = pipe.vae.config.sample_size + elif hasattr(pipe, "movq") and hasattr(pipe.movq.config, "sample_size"): + # Kandinsky + size = pipe.movq.config.sample_size + elif hasattr(pipe, "unet") and hasattr(pipe.unet.config, "sample_size"): + size = pipe.unet.config.sample_size + else: + size = 0 + return size + + +def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): + from modules import lowvram, sd_hijack + checkpoint_info = checkpoint_info or select_checkpoint(op=op) + if checkpoint_info is None: + return + if op == 'model' or op == 'dict': + if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model + return + else: + if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model + return + shared.log.debug(f'Load {op}: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}') + if timer is None: + timer = Timer() + current_checkpoint_info = None + if op == 'model' or op == 'dict': + if model_data.sd_model is not None: + sd_hijack.model_hijack.undo_hijack(model_data.sd_model) + current_checkpoint_info = model_data.sd_model.sd_checkpoint_info + unload_model_weights(op=op) + else: + if model_data.sd_refiner is not None: + sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner) + current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info + unload_model_weights(op=op) + + sd_hijack_inpainting.do_inpainting_hijack() + devices.set_cuda_params() + if already_loaded_state_dict is not None: + state_dict = already_loaded_state_dict + else: + state_dict = get_checkpoint_state_dict(checkpoint_info, timer) + checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info) + if state_dict is None or checkpoint_config is None: + shared.log.error(f"Failed to load checkpooint: {checkpoint_info.filename}") + if current_checkpoint_info is not None: + shared.log.info(f"Restoring previous checkpoint: {current_checkpoint_info.filename}") + load_model(current_checkpoint_info, None) + return + shared.log.debug(f'Model dict loaded: {memory_stats()}') + sd_config = OmegaConf.load(checkpoint_config) + repair_config(sd_config) + timer.record("config") + shared.log.debug(f'Model config loaded: {memory_stats()}') + sd_model = None + stdout = io.StringIO() + if os.environ.get('SD_LDM_DEBUG', None) is not None: + sd_model = instantiate_from_config(sd_config.model) + else: + with contextlib.redirect_stdout(stdout): + """ + try: + clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict + with sd_disable_initialization.DisableInitialization(disable_clip=clip_is_included_into_sd): + sd_model = instantiate_from_config(sd_config.model) + except Exception as e: + shared.log.error(f'LDM: instantiate from config: {e}') + sd_model = instantiate_from_config(sd_config.model) + """ + sd_model = instantiate_from_config(sd_config.model) + for line in stdout.getvalue().splitlines(): + if len(line) > 0: + shared.log.info(f'LDM: {line.strip()}') + shared.log.debug(f"Model created from config: {checkpoint_config}") + sd_model.used_config = checkpoint_config + sd_model.has_accelerate = False + timer.record("create") + ok = load_model_weights(sd_model, checkpoint_info, state_dict, timer) + if not ok: + model_data.sd_model = sd_model + current_checkpoint_info = None + unload_model_weights(op=op) + shared.log.debug(f'Model weights unloaded: {memory_stats()} op={op}') + if op == 'refiner': + # shared.opts.data['sd_model_refiner'] = 'None' + shared.opts.sd_model_refiner = 'None' + return + else: + shared.log.debug(f'Model weights loaded: {memory_stats()}') + timer.record("load") + if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: + lowvram.setup_for_low_vram(sd_model, shared.cmd_opts.medvram) + else: + sd_model.to(devices.device) + timer.record("move") + shared.log.debug(f'Model weights moved: {memory_stats()}') + sd_hijack.model_hijack.hijack(sd_model) + timer.record("hijack") + sd_model.eval() + if op == 'refiner': + model_data.sd_refiner = sd_model + else: + model_data.sd_model = sd_model + sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True) # Reload embeddings after model load as they may or may not fit the model + timer.record("embeddings") + script_callbacks.model_loaded_callback(sd_model) + timer.record("callbacks") + shared.log.info(f"Model loaded in {timer.summary()}") + current_checkpoint_info = None + devices.torch_gc(force=True) + shared.log.info(f'Model load finished: {memory_stats()} cached={len(checkpoints_loaded.keys())}') + + +def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model'): + load_dict = shared.opts.sd_model_dict != model_data.sd_dict + from modules import lowvram, sd_hijack + checkpoint_info = info or select_checkpoint(op=op) # are we selecting model or dictionary + next_checkpoint_info = info or select_checkpoint(op='dict' if load_dict else 'model') if load_dict else None + if checkpoint_info is None: + unload_model_weights(op=op) + return None + orig_state = copy.deepcopy(shared.state) + shared.state = shared_state.State() + shared.state.begin('load') + if load_dict: + shared.log.debug(f'Model dict: existing={sd_model is not None} target={checkpoint_info.filename} info={info}') + else: + model_data.sd_dict = 'None' + shared.log.debug(f'Load model weights: existing={sd_model is not None} target={checkpoint_info.filename} info={info}') + if sd_model is None: + sd_model = model_data.sd_model if op == 'model' or op == 'dict' else model_data.sd_refiner + if sd_model is None: # previous model load failed + current_checkpoint_info = None + else: + current_checkpoint_info = getattr(sd_model, 'sd_checkpoint_info', None) + if current_checkpoint_info is not None and checkpoint_info is not None and current_checkpoint_info.filename == checkpoint_info.filename: + return None + if not getattr(sd_model, 'has_accelerate', False): + if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: + lowvram.send_everything_to_cpu() + else: + sd_model.to(devices.cpu) + if (reuse_dict or shared.opts.model_reuse_dict) and not getattr(sd_model, 'has_accelerate', False): + shared.log.info('Reusing previous model dictionary') + sd_hijack.model_hijack.undo_hijack(sd_model) + else: + unload_model_weights(op=op) + sd_model = None + timer = Timer() + state_dict = get_checkpoint_state_dict(checkpoint_info, timer) + checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info) + timer.record("config") + if sd_model is None or checkpoint_config != getattr(sd_model, 'used_config', None): + sd_model = None + if shared.backend == shared.Backend.ORIGINAL: + load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op) + model_data.sd_dict = shared.opts.sd_model_dict + else: + load_diffuser(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op) + if load_dict and next_checkpoint_info is not None: + model_data.sd_dict = shared.opts.sd_model_dict + shared.opts.data["sd_model_checkpoint"] = next_checkpoint_info.title + reload_model_weights(reuse_dict=True) # ok we loaded dict now lets redo and load model on top of it + shared.state.end() + shared.state = orig_state + # data['sd_model_checkpoint'] + if op == 'model' or op == 'dict': + shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title + return model_data.sd_model + else: + shared.opts.data["sd_model_refiner"] = checkpoint_info.title + return model_data.sd_refiner + + # fallback + shared.log.info(f"Loading using fallback: {op} model={checkpoint_info.title}") + try: + load_model_weights(sd_model, checkpoint_info, state_dict, timer) + except Exception: + shared.log.error("Load model failed: restoring previous") + load_model_weights(sd_model, current_checkpoint_info, None, timer) + finally: + sd_hijack.model_hijack.hijack(sd_model) + timer.record("hijack") + script_callbacks.model_loaded_callback(sd_model) + timer.record("callbacks") + if sd_model is not None and not shared.cmd_opts.lowvram and not shared.cmd_opts.medvram and not getattr(sd_model, 'has_accelerate', False): + sd_model.to(devices.device) + timer.record("device") + shared.state.end() + shared.state = orig_state + shared.log.info(f"Load: {op} time={timer.summary()}") + return sd_model + + +def convert_to_faketensors(tensor): + fake_module = torch._subclasses.fake_tensor.FakeTensorMode(allow_non_fake_inputs=True) # pylint: disable=protected-access + if hasattr(tensor, "weight"): + tensor.weight = torch.nn.Parameter(fake_module.from_tensor(tensor.weight)) + return tensor + + +def disable_offload(sd_model): + from accelerate.hooks import remove_hook_from_module + if not getattr(sd_model, 'has_accelerate', False): + return + for _name, model in sd_model.components.items(): + if not isinstance(model, torch.nn.Module): + continue + remove_hook_from_module(model, recurse=True) + + +def unload_model_weights(op='model', change_from='none'): + if shared.compiled_model_state is not None: + shared.compiled_model_state.compiled_cache.clear() + shared.compiled_model_state.partitioned_modules.clear() + if op == 'model' or op == 'dict': + if model_data.sd_model: + if (shared.backend == shared.Backend.ORIGINAL and change_from != shared.Backend.DIFFUSERS) or change_from == shared.Backend.ORIGINAL: + from modules import sd_hijack + model_data.sd_model.to(devices.cpu) + sd_hijack.model_hijack.undo_hijack(model_data.sd_model) + elif not (shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): + disable_offload(model_data.sd_model) + model_data.sd_model.to('meta') + model_data.sd_model = None + shared.log.debug(f'Unload weights {op}: {memory_stats()}') + else: + if model_data.sd_refiner: + if (shared.backend == shared.Backend.ORIGINAL and change_from != shared.Backend.DIFFUSERS) or change_from == shared.Backend.ORIGINAL: + from modules import sd_hijack + model_data.sd_model.to(devices.cpu) + sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner) + else: + disable_offload(model_data.sd_model) + model_data.sd_refiner.to('meta') + model_data.sd_refiner = None + shared.log.debug(f'Unload weights {op}: {memory_stats()}') + devices.torch_gc(force=True) + + +def apply_token_merging(sd_model, token_merging_ratio=0): + current_token_merging_ratio = getattr(sd_model, 'applied_token_merged_ratio', 0) + if token_merging_ratio is None or current_token_merging_ratio is None or current_token_merging_ratio == token_merging_ratio: + return + try: + if current_token_merging_ratio > 0: + tomesd.remove_patch(sd_model) + except Exception: + pass + if token_merging_ratio > 0: + if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: + shared.log.warning('Token merging not supported with HyperTile for UNet') + return + try: + tomesd.apply_patch( + sd_model, + ratio=token_merging_ratio, + use_rand=False, # can cause issues with some samplers + merge_attn=True, + merge_crossattn=False, + merge_mlp=False + ) + shared.log.info(f'Applying token merging: ratio={token_merging_ratio}') + sd_model.applied_token_merged_ratio = token_merging_ratio + except Exception: + shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') + else: + sd_model.applied_token_merged_ratio = 0 diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 162c2057b..3da8b643e 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -1,690 +1,711 @@ -import os -import html -import csv -import time -from collections import namedtuple -import torch -from tqdm import tqdm -import safetensors.torch -import numpy as np -from PIL import Image, PngImagePlugin -from torch.utils.tensorboard import SummaryWriter -from modules import shared, devices, sd_hijack, processing, sd_models, images, sd_hijack_checkpoint, errors -import modules.textual_inversion.dataset -from modules.textual_inversion.learn_schedule import LearnRateScheduler -from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay -from modules.textual_inversion.ti_logging import save_settings_to_file -from modules.files_cache import directory_files, extension_filter, directory_mtime - -TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"]) -textual_inversion_templates = {} - - -def list_textual_inversion_templates(): - textual_inversion_templates.clear() - for root, _dirs, fns in os.walk(shared.opts.embeddings_templates_dir): - for fn in fns: - path = os.path.join(root, fn) - textual_inversion_templates[fn] = TextualInversionTemplate(fn, path) - return textual_inversion_templates - - -class Embedding: - def __init__(self, vec, name, filename=None, step=None): - self.vec = vec - self.name = name - self.tag = name - self.step = step - self.filename = filename - self.basename = os.path.relpath(filename, shared.opts.embeddings_dir) if filename is not None else None - self.shape = None - self.vectors = 0 - self.cached_checksum = None - self.sd_checkpoint = None - self.sd_checkpoint_name = None - self.optimizer_state_dict = None - - def save(self, filename): - embedding_data = { - "string_to_token": {"*": 265}, - "string_to_param": {"*": self.vec}, - "name": self.name, - "step": self.step, - "sd_checkpoint": self.sd_checkpoint, - "sd_checkpoint_name": self.sd_checkpoint_name, - } - torch.save(embedding_data, filename) - if shared.opts.save_optimizer_state and self.optimizer_state_dict is not None: - optimizer_saved_dict = { - 'hash': self.checksum(), - 'optimizer_state_dict': self.optimizer_state_dict, - } - torch.save(optimizer_saved_dict, f"{filename}.optim") - - def checksum(self): - if self.cached_checksum is not None: - return self.cached_checksum - def const_hash(a): - r = 0 - for v in a: - r = (r * 281 ^ int(v) * 997) & 0xFFFFFFFF - return r - self.cached_checksum = f'{const_hash(self.vec.reshape(-1) * 100) & 0xffff:04x}' - return self.cached_checksum - - -class DirWithTextualInversionEmbeddings: - def __init__(self, path): - self.path = path - self.mtime = None - - def has_changed(self): - if not os.path.isdir(self.path): - return False - return directory_mtime(self.path) != self.mtime - - def update(self): - if not os.path.isdir(self.path): - return - self.mtime = directory_mtime(self.path) - - -class EmbeddingDatabase: - def __init__(self): - self.ids_lookup = {} - self.word_embeddings = {} - self.skipped_embeddings = {} - self.expected_shape = -1 - self.embedding_dirs = {} - self.previously_displayed_embeddings = () - self.embeddings_used = [] - - def add_embedding_dir(self, path): - self.embedding_dirs[path] = DirWithTextualInversionEmbeddings(path) - - def clear_embedding_dirs(self): - self.embedding_dirs.clear() - - def register_embedding(self, embedding, model): - self.word_embeddings[embedding.name] = embedding - if hasattr(model, 'cond_stage_model'): - ids = model.cond_stage_model.tokenize([embedding.name])[0] - elif hasattr(model, 'tokenizer'): - ids = model.tokenizer.convert_tokens_to_ids(embedding.name) - if type(ids) != list: - ids = [ids] - first_id = ids[0] - if first_id not in self.ids_lookup: - self.ids_lookup[first_id] = [] - self.ids_lookup[first_id] = sorted(self.ids_lookup[first_id] + [(ids, embedding)], key=lambda x: len(x[0]), reverse=True) - return embedding - - def get_expected_shape(self): - if shared.backend == shared.Backend.DIFFUSERS: - return 0 - if shared.sd_model is None: - shared.log.error('Model not loaded') - return 0 - vec = shared.sd_model.cond_stage_model.encode_embedding_init_text(",", 1) - return vec.shape[1] - - def load_diffusers_embedding(self, filename: str, path: str): - if shared.sd_model is None: - return - fn, ext = os.path.splitext(filename) - if ext.lower() != ".pt" and ext.lower() != ".safetensors": - return - pipe = shared.sd_model - name = os.path.basename(fn) - embedding = Embedding(vec=None, name=name, filename=path) - if not hasattr(pipe, "tokenizer") or not hasattr(pipe, 'text_encoder'): - self.skipped_embeddings[name] = embedding - return - try: - is_xl = hasattr(pipe, 'text_encoder_2') - try: - if not is_xl: # only use for sd15/sd21 - pipe.load_textual_inversion(path, token=name, cache_dir=shared.opts.diffusers_dir, local_files_only=True) - self.register_embedding(embedding, shared.sd_model) - except Exception: - pass - is_loaded = pipe.tokenizer.convert_tokens_to_ids(name) - if type(is_loaded) != list: - is_loaded = [is_loaded] - is_loaded = is_loaded[0] > 49407 - if is_loaded: - self.register_embedding(embedding, shared.sd_model) - else: - embeddings_dict = {} - if ext.lower() in ['.safetensors']: - with safetensors.torch.safe_open(path, framework="pt") as f: - for k in f.keys(): - embeddings_dict[k] = f.get_tensor(k) - else: - raise NotImplementedError - """ - # alternatively could disable load_textual_inversion and load everything here - elif ext.lower() in ['.pt', '.bin']: - data = torch.load(path, map_location="cpu") - embedding.tag = data.get('name', None) - embedding.step = data.get('step', None) - embedding.sd_checkpoint = data.get('sd_checkpoint', None) - embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None) - param_dict = data.get('string_to_param', None) - embeddings_dict['clip_l'] = [] - for tokens in param_dict.values(): - for vec in tokens: - embeddings_dict['clip_l'].append(vec) - """ - clip_l = pipe.text_encoder if hasattr(pipe, 'text_encoder') else None - clip_g = pipe.text_encoder_2 if hasattr(pipe, 'text_encoder_2') else None - is_sd = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is None and 'clip_g' not in embeddings_dict - is_xl = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is not None and 'clip_g' in embeddings_dict - tokens = [] - for i in range(len(embeddings_dict["clip_l"])): - if (is_sd or is_xl) and (len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i])): - tokens.append(name if i == 0 else f"{name}_{i}") - num_added = pipe.tokenizer.add_tokens(tokens) - if num_added > 0: - token_ids = pipe.tokenizer.convert_tokens_to_ids(tokens) - if is_sd: # only used for sd15 if load_textual_inversion failed and format is safetensors - clip_l.resize_token_embeddings(len(pipe.tokenizer)) - for i in range(len(token_ids)): - clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i] - elif is_xl: - pipe.tokenizer_2.add_tokens(tokens) - clip_l.resize_token_embeddings(len(pipe.tokenizer)) - clip_g.resize_token_embeddings(len(pipe.tokenizer)) - for i in range(len(token_ids)): - clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i] - clip_g.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_g"][i] - self.register_embedding(embedding, shared.sd_model) - else: - raise NotImplementedError - except Exception: - self.skipped_embeddings[name] = embedding - - def load_from_file(self, path, filename): - name, ext = os.path.splitext(filename) - ext = ext.upper() - if shared.backend == shared.Backend.DIFFUSERS: - self.load_diffusers_embedding(filename, path) - return - - if ext in ['.PNG', '.WEBP', '.JXL', '.AVIF']: - if '.preview' in filename.lower(): - return - embed_image = Image.open(path) - if hasattr(embed_image, 'text') and 'sd-ti-embedding' in embed_image.text: - data = embedding_from_b64(embed_image.text['sd-ti-embedding']) - else: - data = extract_image_data_embed(embed_image) - if not data: # if data is None, means this is not an embeding, just a preview image - return - elif ext in ['.BIN', '.PT']: - data = torch.load(path, map_location="cpu") - elif ext in ['.SAFETENSORS']: - data = safetensors.torch.load_file(path, device="cpu") - else: - return - - # textual inversion embeddings - if 'string_to_param' in data: - param_dict = data['string_to_param'] - param_dict = getattr(param_dict, '_parameters', param_dict) # fix for torch 1.12.1 loading saved file from torch 1.11 - assert len(param_dict) == 1, 'embedding file has multiple terms in it' - emb = next(iter(param_dict.items()))[1] - # diffuser concepts - elif type(data) == dict and type(next(iter(data.values()))) == torch.Tensor: - if len(data.keys()) != 1: - self.skipped_embeddings[name] = Embedding(None, name=name, filename=path) - return - emb = next(iter(data.values())) - if len(emb.shape) == 1: - emb = emb.unsqueeze(0) - else: - raise RuntimeError(f"Couldn't identify {filename} as textual inversion embedding") - - vec = emb.detach().to(devices.device, dtype=torch.float32) - # name = data.get('name', name) - embedding = Embedding(vec=vec, name=name, filename=path) - embedding.tag = data.get('name', None) - embedding.step = data.get('step', None) - embedding.sd_checkpoint = data.get('sd_checkpoint', None) - embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None) - embedding.vectors = vec.shape[0] - embedding.shape = vec.shape[-1] - if self.expected_shape == -1 or self.expected_shape == embedding.shape: - self.register_embedding(embedding, shared.sd_model) - else: - self.skipped_embeddings[name] = embedding - - def load_from_dir(self, embdir): - if sd_models.model_data.sd_model is None: - shared.log.info('Skipping embeddings load: model not loaded') - return - if not os.path.isdir(embdir.path): - return - is_ext = extension_filter(['.PNG', '.WEBP', '.JXL', '.AVIF', '.BIN', '.PT', '.SAFETENSORS']) - is_not_preview = lambda fp: not next(iter(os.path.splitext(fp))).upper().endswith('.PREVIEW') # pylint: disable=unnecessary-lambda-assignment - for file_path in [*filter(lambda fp: is_ext(fp) and is_not_preview(fp), directory_files(embdir.path))]: - try: - if os.stat(file_path).st_size == 0: - continue - fn = os.path.basename(file_path) - self.load_from_file(file_path, fn) - except Exception as e: - errors.display(e, f'embedding load {fn}') - continue - - def load_textual_inversion_embeddings(self, force_reload=False): - if shared.sd_model is None: - return - t0 = time.time() - if not force_reload: - need_reload = False - for embdir in self.embedding_dirs.values(): - if embdir.has_changed(): - need_reload = True - break - if not need_reload: - return - self.ids_lookup.clear() - self.word_embeddings.clear() - self.skipped_embeddings.clear() - self.embeddings_used.clear() - self.expected_shape = self.get_expected_shape() - for embdir in self.embedding_dirs.values(): - self.load_from_dir(embdir) - embdir.update() - - # re-sort word_embeddings because load_from_dir may not load in alphabetic order. - # using a temporary copy so we don't reinitialize self.word_embeddings in case other objects have a reference to it. - sorted_word_embeddings = {e.name: e for e in sorted(self.word_embeddings.values(), key=lambda e: e.name.lower())} - self.word_embeddings.clear() - self.word_embeddings.update(sorted_word_embeddings) - - displayed_embeddings = (tuple(self.word_embeddings.keys()), tuple(self.skipped_embeddings.keys())) - if self.previously_displayed_embeddings != displayed_embeddings: - self.previously_displayed_embeddings = displayed_embeddings - t1 = time.time() - shared.log.info(f"Load embeddings: loaded={len(self.word_embeddings)} skipped={len(self.skipped_embeddings)} time={t1-t0:.2f}") - - - def find_embedding_at_position(self, tokens, offset): - token = tokens[offset] - possible_matches = self.ids_lookup.get(token, None) - if possible_matches is None: - return None, None - for ids, embedding in possible_matches: - if tokens[offset:offset + len(ids)] == ids: - return embedding, len(ids) - return None, None - - -def create_embedding(name, num_vectors_per_token, overwrite_old, init_text='*'): - cond_model = shared.sd_model.cond_stage_model - with devices.autocast(): - cond_model([""]) # will send cond model to GPU if lowvram/medvram is active - #cond_model expects at least some text, so we provide '*' as backup. - embedded = cond_model.encode_embedding_init_text(init_text or '*', num_vectors_per_token) - vec = torch.zeros((num_vectors_per_token, embedded.shape[1]), device=devices.device) - #Only copy if we provided an init_text, otherwise keep vectors as zeros - if init_text: - for i in range(num_vectors_per_token): - vec[i] = embedded[i * int(embedded.shape[0]) // num_vectors_per_token] - # Remove illegal characters from name. - name = "".join( x for x in name if (x.isalnum() or x in "._- ")) - fn = os.path.join(shared.opts.embeddings_dir, f"{name}.pt") - if not overwrite_old and os.path.exists(fn): - shared.log.warning(f"Embedding already exists: {fn}") - else: - embedding = Embedding(vec=vec, name=name, filename=fn) - embedding.step = 0 - embedding.save(fn) - shared.log.info(f'Created embedding: {fn} vectors {num_vectors_per_token} init {init_text}') - return fn - - -def write_loss(log_directory, filename, step, epoch_len, values): - if shared.opts.training_write_csv_every == 0: - return - if step % shared.opts.training_write_csv_every != 0: - return - write_csv_header = False if os.path.exists(os.path.join(log_directory, filename)) else True - with open(os.path.join(log_directory, filename), "a+", newline='', encoding='utf-8') as fout: - csv_writer = csv.DictWriter(fout, fieldnames=["step", "epoch", "epoch_step", *(values.keys())]) - if write_csv_header: - csv_writer.writeheader() - epoch = (step - 1) // epoch_len - epoch_step = (step - 1) % epoch_len - csv_writer.writerow({ - "step": step, - "epoch": epoch, - "epoch_step": epoch_step, - **values, - }) - - -def tensorboard_setup(log_directory): - os.makedirs(os.path.join(log_directory, "tensorboard"), exist_ok=True) - return SummaryWriter( - log_dir=os.path.join(log_directory, "tensorboard"), - flush_secs=shared.opts.training_tensorboard_flush_every) - - -def tensorboard_add(tensorboard_writer, loss, global_step, step, learn_rate, epoch_num): - tensorboard_add_scaler(tensorboard_writer, "Loss/train", loss, global_step) - tensorboard_add_scaler(tensorboard_writer, f"Loss/train/epoch-{epoch_num}", loss, step) - tensorboard_add_scaler(tensorboard_writer, "Learn rate/train", learn_rate, global_step) - tensorboard_add_scaler(tensorboard_writer, f"Learn rate/train/epoch-{epoch_num}", learn_rate, step) - - -def tensorboard_add_scaler(tensorboard_writer, tag, value, step): - tensorboard_writer.add_scalar(tag=tag, scalar_value=value, global_step=step) - - -def tensorboard_add_image(tensorboard_writer, tag, pil_image, step): - # Convert a pil image to a torch tensor - img_tensor = torch.as_tensor(np.array(pil_image, copy=True)) - img_tensor = img_tensor.view(pil_image.size[1], pil_image.size[0], len(pil_image.getbands())) - img_tensor = img_tensor.permute((2, 0, 1)) - tensorboard_writer.add_image(tag, img_tensor, global_step=step) - - -def validate_train_inputs(model_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_model_every, create_image_every, name="embedding"): - assert model_name, f"{name} not selected" - assert learn_rate, "Learning rate is empty or 0" - assert isinstance(batch_size, int), "Batch size must be integer" - assert batch_size > 0, "Batch size must be positive" - assert isinstance(gradient_step, int), "Gradient accumulation step must be integer" - assert gradient_step > 0, "Gradient accumulation step must be positive" - assert data_root, "Dataset directory is empty" - assert os.path.isdir(data_root), "Dataset directory doesn't exist" - assert os.listdir(data_root), "Dataset directory is empty" - assert template_filename, "Prompt template file not selected" - assert template_file, f"Prompt template file {template_filename} not found" - assert os.path.isfile(template_file.path), f"Prompt template file {template_filename} doesn't exist" - assert steps, "Max steps is empty or 0" - assert isinstance(steps, int), "Max steps must be integer" - assert steps > 0, "Max steps must be positive" - assert isinstance(save_model_every, int), "Save {name} must be integer" - assert save_model_every >= 0, "Save {name} must be positive or 0" - assert isinstance(create_image_every, int), "Create image must be integer" - assert create_image_every >= 0, "Create image must be positive or 0" - - -def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument - - shared.log.debug(f'train_embedding: embedding_name={embedding_name}|learn_rate={learn_rate}|batch_size={batch_size}|gradient_step={gradient_step}|data_root={data_root}|log_directory={log_directory}|training_width={training_width}|training_height={training_height}|varsize={varsize}|steps={steps}|clip_grad_mode={clip_grad_mode}|clip_grad_value={clip_grad_value}|shuffle_tags={shuffle_tags}|tag_drop_out={tag_drop_out}|latent_sampling_method={latent_sampling_method}|use_weight={use_weight}|create_image_every={create_image_every}|save_embedding_every={save_embedding_every}|template_filename={template_filename}|save_image_with_stored_embedding={save_image_with_stored_embedding}|preview_from_txt2img={preview_from_txt2img}|preview_prompt={preview_prompt}|preview_negative_prompt={preview_negative_prompt}|preview_steps={preview_steps}|preview_sampler_index={preview_sampler_index}|preview_cfg_scale={preview_cfg_scale}|preview_seed={preview_seed}|preview_width={preview_width}|preview_height={preview_height}') - save_embedding_every = save_embedding_every or 0 - create_image_every = create_image_every or 0 - template_file = textual_inversion_templates.get(template_filename, None) - validate_train_inputs(embedding_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_embedding_every, create_image_every, name="embedding") - if log_directory is None or log_directory == '': - log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}" - template_file = template_file.path - - shared.state.job = "train" - shared.state.textinfo = "Initializing textual inversion training..." - shared.state.job_count = steps - - filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt') - - if log_directory == '': - log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}" - log_directory = os.path.join(log_directory, embedding_name) - unload = shared.opts.unload_models_when_training - - if save_embedding_every > 0: - embedding_dir = os.path.join(log_directory, "embeddings") - os.makedirs(embedding_dir, exist_ok=True) - else: - embedding_dir = None - - if create_image_every > 0: - images_dir = os.path.join(log_directory, "images") - os.makedirs(images_dir, exist_ok=True) - else: - images_dir = None - - if create_image_every > 0 and save_image_with_stored_embedding: - images_embeds_dir = os.path.join(log_directory, "image_embeddings") - os.makedirs(images_embeds_dir, exist_ok=True) - else: - images_embeds_dir = None - - hijack = sd_hijack.model_hijack - embedding = hijack.embedding_db.word_embeddings[embedding_name] - checkpoint = sd_models.select_checkpoint() - initial_step = embedding.step or 0 - if initial_step >= steps: - shared.state.textinfo = "Model has already been trained beyond specified max steps" - return embedding, filename - scheduler = LearnRateScheduler(learn_rate, steps, initial_step) - clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else \ - torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else \ - None - if clip_grad: - clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False) - # dataset loading may take a while, so input validations and early returns should be done before this - shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..." - old_parallel_processing_allowed = shared.parallel_processing_allowed - - if shared.opts.training_enable_tensorboard: - tensorboard_writer = tensorboard_setup(log_directory) - - pin_memory = shared.opts.pin_memory - # init dataset - ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=embedding_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight) - - if shared.opts.save_training_settings_to_txt: - save_settings_to_file(log_directory, {**dict(model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds), num_vectors_per_token=len(embedding.vec)), **locals()}) - latent_sampling_method = ds.latent_sampling_method - # init dataloader - dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory) - if unload: - shared.parallel_processing_allowed = False - shared.sd_model.first_stage_model.to(devices.cpu) - - embedding.vec.requires_grad = True - optimizer = torch.optim.AdamW([embedding.vec], lr=scheduler.learn_rate, weight_decay=0.0) - if shared.opts.save_optimizer_state: - optimizer_state_dict = None - if os.path.exists(f"{filename}.optim"): - optimizer_saved_dict = torch.load(f"{filename}.optim", map_location='cpu') - if embedding.checksum() == optimizer_saved_dict.get('hash', None): - optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None) - if optimizer_state_dict is not None: - optimizer.load_state_dict(optimizer_state_dict) - shared.log.info("Load existing optimizer from checkpoint") - else: - shared.log.info("No saved optimizer exists in checkpoint") - - scaler = torch.cuda.amp.GradScaler() - - batch_size = ds.batch_size - gradient_step = ds.gradient_step - # n steps = batch_size * gradient_step * n image processed - steps_per_epoch = len(ds) // batch_size // gradient_step - max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step - loss_step = 0 - _loss_step = 0 #internal - last_saved_file = "" - last_saved_image = "" - forced_filename = "" - embedding_yet_to_be_embedded = False - is_training_inpainting_model = shared.sd_model.model.conditioning_key in {'hybrid', 'concat'} - img_c = None - - pbar = tqdm(total=steps - initial_step) - try: - sd_hijack_checkpoint.add() - for _i in range((steps-initial_step) * gradient_step): - if scheduler.finished: - break - if shared.state.interrupted: - break - for j, batch in enumerate(dl): - # works as a drop_last=True for gradient accumulation - if j == max_steps_per_epoch: - break - scheduler.apply(optimizer, embedding.step) - if scheduler.finished: - break - if shared.state.interrupted: - break - if clip_grad: - clip_grad_sched.step(embedding.step) - with devices.autocast(): - x = batch.latent_sample.to(devices.device, non_blocking=pin_memory) - if use_weight: - w = batch.weight.to(devices.device, non_blocking=pin_memory) - c = shared.sd_model.cond_stage_model(batch.cond_text) - if is_training_inpainting_model: - if img_c is None: - img_c = processing.txt2img_image_conditioning(shared.sd_model, c, training_width, training_height) - cond = {"c_concat": [img_c], "c_crossattn": [c]} - else: - cond = c - if use_weight: - loss = shared.sd_model.weighted_forward(x, cond, w)[0] / gradient_step - del w - else: - loss = shared.sd_model.forward(x, cond)[0] / gradient_step - del x - _loss_step += loss.item() - - scaler.scale(loss).backward() - # go back until we reach gradient accumulation steps - if (j + 1) % gradient_step != 0: - continue - if clip_grad: - clip_grad(embedding.vec, clip_grad_sched.learn_rate) - - scaler.step(optimizer) - scaler.update() - embedding.step += 1 - pbar.update() - optimizer.zero_grad(set_to_none=True) - loss_step = _loss_step - _loss_step = 0 - steps_done = embedding.step + 1 - epoch_num = embedding.step // steps_per_epoch - - description = f"Training textual inversion step {embedding.step} loss: {loss_step:.5f} lr: {scheduler.learn_rate:.5f}" - pbar.set_description(description) - if embedding_dir is not None and steps_done % save_embedding_every == 0: - # Before saving, change name to match current checkpoint. - embedding_name_every = f'{embedding_name}-{steps_done}' - last_saved_file = os.path.join(embedding_dir, f'{embedding_name_every}.pt') - save_embedding(embedding, optimizer, checkpoint, embedding_name_every, last_saved_file, remove_cached_checksum=True) - embedding_yet_to_be_embedded = True - - write_loss(log_directory, f"{embedding_name}.csv", embedding.step, steps_per_epoch, { "loss": f"{loss_step:.7f}", "learn_rate": scheduler.learn_rate }) - - if images_dir is not None and steps_done % create_image_every == 0: - forced_filename = f'{embedding_name}-{steps_done}' - last_saved_image = os.path.join(images_dir, forced_filename) - shared.sd_model.first_stage_model.to(devices.device) - - p = processing.StableDiffusionProcessingTxt2Img( - sd_model=shared.sd_model, - do_not_save_grid=True, - do_not_save_samples=True, - do_not_reload_embeddings=True, - ) - - if preview_from_txt2img: - p.prompt = preview_prompt - p.negative_prompt = preview_negative_prompt - p.steps = preview_steps - p.sampler_name = processing.get_sampler_name(preview_sampler_index) - p.cfg_scale = preview_cfg_scale - p.seed = preview_seed - p.width = preview_width - p.height = preview_height - else: - p.prompt = batch.cond_text[0] - p.steps = 20 - p.width = training_width - p.height = training_height - - preview_text = p.prompt - processed = processing.process_images(p) - image = processed.images[0] if len(processed.images) > 0 else None - - if unload: - shared.sd_model.first_stage_model.to(devices.cpu) - - if image is not None: - shared.state.assign_current_image(image) - last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) - last_saved_image += f", prompt: {preview_text}" - if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images: - tensorboard_add_image(tensorboard_writer, f"Validation at epoch {epoch_num}", image, embedding.step) - - if save_image_with_stored_embedding and os.path.exists(last_saved_file) and embedding_yet_to_be_embedded: - last_saved_image_chunks = os.path.join(images_embeds_dir, f'{embedding_name}-{steps_done}.png') - info = PngImagePlugin.PngInfo() - data = torch.load(last_saved_file) - info.add_text("sd-ti-embedding", embedding_to_b64(data)) - title = f"<{data.get('name', '???')}>" - try: - vectorSize = list(data['string_to_param'].values())[0].shape[0] - except Exception: - vectorSize = '?' - checkpoint = sd_models.select_checkpoint() - footer_left = checkpoint.model_name - footer_mid = f'[{checkpoint.shorthash}]' - footer_right = f'{vectorSize}v {steps_done}s' - captioned_image = caption_image_overlay(image, title, footer_left, footer_mid, footer_right) - captioned_image = insert_image_data_embed(captioned_image, data) - captioned_image.save(last_saved_image_chunks, "PNG", pnginfo=info) - embedding_yet_to_be_embedded = False - - last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) - last_saved_image += f", prompt: {preview_text}" - - shared.state.job_no = embedding.step - shared.state.textinfo = f""" -

-Loss: {loss_step:.7f}
-Step: {steps_done}
-Last prompt: {html.escape(batch.cond_text[0])}
-Last saved embedding: {html.escape(last_saved_file)}
-Last saved image: {html.escape(last_saved_image)}
-

-""" - filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt') - save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True) - except Exception as e: - errors.display(e, 'embedding train') - finally: - pbar.leave = False - pbar.close() - shared.sd_model.first_stage_model.to(devices.device) - shared.parallel_processing_allowed = old_parallel_processing_allowed - sd_hijack_checkpoint.remove() - return embedding, filename - - -def save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True): - old_embedding_name = embedding.name - old_sd_checkpoint = embedding.sd_checkpoint if hasattr(embedding, "sd_checkpoint") else None - old_sd_checkpoint_name = embedding.sd_checkpoint_name if hasattr(embedding, "sd_checkpoint_name") else None - old_cached_checksum = embedding.cached_checksum if hasattr(embedding, "cached_checksum") else None - try: - embedding.sd_checkpoint = checkpoint.shorthash - embedding.sd_checkpoint_name = checkpoint.model_name - if remove_cached_checksum: - embedding.cached_checksum = None - embedding.name = embedding_name - embedding.optimizer_state_dict = optimizer.state_dict() - embedding.save(filename) - except Exception: - embedding.sd_checkpoint = old_sd_checkpoint - embedding.sd_checkpoint_name = old_sd_checkpoint_name - embedding.name = old_embedding_name - embedding.cached_checksum = old_cached_checksum - raise +import csv +import html +import os +import time +from collections import namedtuple + +import numpy as np +import safetensors.torch +import torch +from PIL import Image, PngImagePlugin +from torch.utils.tensorboard import SummaryWriter +from tqdm import tqdm + +import modules.textual_inversion.dataset +from modules import ( + devices, + errors, + images, + processing, + sd_hijack, + sd_hijack_checkpoint, + sd_models, + shared, +) +from modules.files_cache import ( + directory_files, + directory_mtime, + extension_filter, +) +from modules.textual_inversion.image_embedding import ( + caption_image_overlay, + embedding_from_b64, + embedding_to_b64, + extract_image_data_embed, + insert_image_data_embed, +) +from modules.textual_inversion.learn_schedule import LearnRateScheduler +from modules.textual_inversion.ti_logging import save_settings_to_file + +TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"]) +textual_inversion_templates = {} + + +def list_textual_inversion_templates(): + textual_inversion_templates.clear() + for root, _dirs, fns in os.walk(shared.opts.embeddings_templates_dir): + for fn in fns: + path = os.path.join(root, fn) + textual_inversion_templates[fn] = TextualInversionTemplate(fn, path) + return textual_inversion_templates + + +class Embedding: + def __init__(self, vec, name, filename=None, step=None): + self.vec = vec + self.name = name + self.tag = name + self.step = step + self.filename = filename + self.basename = os.path.relpath(filename, shared.opts.embeddings_dir) if filename is not None else None + self.shape = None + self.vectors = 0 + self.cached_checksum = None + self.sd_checkpoint = None + self.sd_checkpoint_name = None + self.optimizer_state_dict = None + + def save(self, filename): + embedding_data = { + "string_to_token": {"*": 265}, + "string_to_param": {"*": self.vec}, + "name": self.name, + "step": self.step, + "sd_checkpoint": self.sd_checkpoint, + "sd_checkpoint_name": self.sd_checkpoint_name, + } + torch.save(embedding_data, filename) + if shared.opts.save_optimizer_state and self.optimizer_state_dict is not None: + optimizer_saved_dict = { + 'hash': self.checksum(), + 'optimizer_state_dict': self.optimizer_state_dict, + } + torch.save(optimizer_saved_dict, f"{filename}.optim") + + def checksum(self): + if self.cached_checksum is not None: + return self.cached_checksum + def const_hash(a): + r = 0 + for v in a: + r = (r * 281 ^ int(v) * 997) & 0xFFFFFFFF + return r + self.cached_checksum = f'{const_hash(self.vec.reshape(-1) * 100) & 0xffff:04x}' + return self.cached_checksum + + +class DirWithTextualInversionEmbeddings: + def __init__(self, path): + self.path = path + self.mtime = None + + def has_changed(self): + if not os.path.isdir(self.path): + return False + return directory_mtime(self.path) != self.mtime + + def update(self): + if not os.path.isdir(self.path): + return + self.mtime = directory_mtime(self.path) + + +class EmbeddingDatabase: + def __init__(self): + self.ids_lookup = {} + self.word_embeddings = {} + self.skipped_embeddings = {} + self.expected_shape = -1 + self.embedding_dirs = {} + self.previously_displayed_embeddings = () + self.embeddings_used = [] + + def add_embedding_dir(self, path): + self.embedding_dirs[path] = DirWithTextualInversionEmbeddings(path) + + def clear_embedding_dirs(self): + self.embedding_dirs.clear() + + def register_embedding(self, embedding, model): + self.word_embeddings[embedding.name] = embedding + if hasattr(model, 'cond_stage_model'): + ids = model.cond_stage_model.tokenize([embedding.name])[0] + elif hasattr(model, 'tokenizer'): + ids = model.tokenizer.convert_tokens_to_ids(embedding.name) + if type(ids) != list: + ids = [ids] + first_id = ids[0] + if first_id not in self.ids_lookup: + self.ids_lookup[first_id] = [] + self.ids_lookup[first_id] = sorted(self.ids_lookup[first_id] + [(ids, embedding)], key=lambda x: len(x[0]), reverse=True) + return embedding + + def get_expected_shape(self): + if shared.backend == shared.Backend.DIFFUSERS: + return 0 + if shared.sd_model is None: + shared.log.error('Model not loaded') + return 0 + vec = shared.sd_model.cond_stage_model.encode_embedding_init_text(",", 1) + return vec.shape[1] + + def load_diffusers_embedding(self, filename: str, path: str): + if shared.sd_model is None: + return + fn, ext = os.path.splitext(filename) + if ext.lower() != ".pt" and ext.lower() != ".safetensors": + return + pipe = shared.sd_model + name = os.path.basename(fn) + embedding = Embedding(vec=None, name=name, filename=path) + if not hasattr(pipe, "tokenizer") or not hasattr(pipe, 'text_encoder'): + self.skipped_embeddings[name] = embedding + return + try: + is_xl = hasattr(pipe, 'text_encoder_2') + try: + if not is_xl: # only use for sd15/sd21 + pipe.load_textual_inversion(path, token=name, cache_dir=shared.opts.diffusers_dir, local_files_only=True) + self.register_embedding(embedding, shared.sd_model) + except Exception: + pass + is_loaded = pipe.tokenizer.convert_tokens_to_ids(name) + if type(is_loaded) != list: + is_loaded = [is_loaded] + is_loaded = is_loaded[0] > 49407 + if is_loaded: + self.register_embedding(embedding, shared.sd_model) + else: + embeddings_dict = {} + if ext.lower() in ['.safetensors']: + with safetensors.torch.safe_open(path, framework="pt") as f: + for k in f.keys(): + embeddings_dict[k] = f.get_tensor(k) + else: + raise NotImplementedError + """ + # alternatively could disable load_textual_inversion and load everything here + elif ext.lower() in ['.pt', '.bin']: + data = torch.load(path, map_location="cpu") + embedding.tag = data.get('name', None) + embedding.step = data.get('step', None) + embedding.sd_checkpoint = data.get('sd_checkpoint', None) + embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None) + param_dict = data.get('string_to_param', None) + embeddings_dict['clip_l'] = [] + for tokens in param_dict.values(): + for vec in tokens: + embeddings_dict['clip_l'].append(vec) + """ + clip_l = pipe.text_encoder if hasattr(pipe, 'text_encoder') else None + clip_g = pipe.text_encoder_2 if hasattr(pipe, 'text_encoder_2') else None + is_sd = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is None and 'clip_g' not in embeddings_dict + is_xl = clip_l is not None and 'clip_l' in embeddings_dict and clip_g is not None and 'clip_g' in embeddings_dict + tokens = [] + for i in range(len(embeddings_dict["clip_l"])): + if (is_sd or is_xl) and (len(clip_l.get_input_embeddings().weight.data[0]) == len(embeddings_dict["clip_l"][i])): + tokens.append(name if i == 0 else f"{name}_{i}") + num_added = pipe.tokenizer.add_tokens(tokens) + if num_added > 0: + token_ids = pipe.tokenizer.convert_tokens_to_ids(tokens) + if is_sd: # only used for sd15 if load_textual_inversion failed and format is safetensors + clip_l.resize_token_embeddings(len(pipe.tokenizer)) + for i in range(len(token_ids)): + clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i] + elif is_xl: + pipe.tokenizer_2.add_tokens(tokens) + clip_l.resize_token_embeddings(len(pipe.tokenizer)) + clip_g.resize_token_embeddings(len(pipe.tokenizer)) + for i in range(len(token_ids)): + clip_l.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_l"][i] + clip_g.get_input_embeddings().weight.data[token_ids[i]] = embeddings_dict["clip_g"][i] + self.register_embedding(embedding, shared.sd_model) + else: + raise NotImplementedError + except Exception: + self.skipped_embeddings[name] = embedding + + def load_from_file(self, path, filename): + name, ext = os.path.splitext(filename) + ext = ext.upper() + if shared.backend == shared.Backend.DIFFUSERS: + self.load_diffusers_embedding(filename, path) + return + + if ext in ['.PNG', '.WEBP', '.JXL', '.AVIF']: + if '.preview' in filename.lower(): + return + embed_image = Image.open(path) + if hasattr(embed_image, 'text') and 'sd-ti-embedding' in embed_image.text: + data = embedding_from_b64(embed_image.text['sd-ti-embedding']) + else: + data = extract_image_data_embed(embed_image) + if not data: # if data is None, means this is not an embeding, just a preview image + return + elif ext in ['.BIN', '.PT']: + data = torch.load(path, map_location="cpu") + elif ext in ['.SAFETENSORS']: + data = safetensors.torch.load_file(path, device="cpu") + else: + return + + # textual inversion embeddings + if 'string_to_param' in data: + param_dict = data['string_to_param'] + param_dict = getattr(param_dict, '_parameters', param_dict) # fix for torch 1.12.1 loading saved file from torch 1.11 + assert len(param_dict) == 1, 'embedding file has multiple terms in it' + emb = next(iter(param_dict.items()))[1] + # diffuser concepts + elif type(data) == dict and type(next(iter(data.values()))) == torch.Tensor: + if len(data.keys()) != 1: + self.skipped_embeddings[name] = Embedding(None, name=name, filename=path) + return + emb = next(iter(data.values())) + if len(emb.shape) == 1: + emb = emb.unsqueeze(0) + else: + raise RuntimeError(f"Couldn't identify {filename} as textual inversion embedding") + + vec = emb.detach().to(devices.device, dtype=torch.float32) + # name = data.get('name', name) + embedding = Embedding(vec=vec, name=name, filename=path) + embedding.tag = data.get('name', None) + embedding.step = data.get('step', None) + embedding.sd_checkpoint = data.get('sd_checkpoint', None) + embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None) + embedding.vectors = vec.shape[0] + embedding.shape = vec.shape[-1] + if self.expected_shape == -1 or self.expected_shape == embedding.shape: + self.register_embedding(embedding, shared.sd_model) + else: + self.skipped_embeddings[name] = embedding + + def load_from_dir(self, embdir): + if sd_models.model_data.sd_model is None: + shared.log.info('Skipping embeddings load: model not loaded') + return + if not os.path.isdir(embdir.path): + return + is_ext = extension_filter(['.PNG', '.WEBP', '.JXL', '.AVIF', '.BIN', '.PT', '.SAFETENSORS']) + is_not_preview = lambda fp: not next(iter(os.path.splitext(fp))).upper().endswith('.PREVIEW') # pylint: disable=unnecessary-lambda-assignment + for file_path in [*filter(lambda fp: is_ext(fp) and is_not_preview(fp), directory_files(embdir.path))]: + try: + if os.stat(file_path).st_size == 0: + continue + fn = os.path.basename(file_path) + self.load_from_file(file_path, fn) + except Exception as e: + errors.display(e, f'embedding load {fn}') + continue + + def load_textual_inversion_embeddings(self, force_reload=False): + if shared.sd_model is None: + return + t0 = time.time() + if not force_reload: + need_reload = False + for embdir in self.embedding_dirs.values(): + if embdir.has_changed(): + need_reload = True + break + if not need_reload: + return + self.ids_lookup.clear() + self.word_embeddings.clear() + self.skipped_embeddings.clear() + self.embeddings_used.clear() + self.expected_shape = self.get_expected_shape() + for embdir in self.embedding_dirs.values(): + self.load_from_dir(embdir) + embdir.update() + + # re-sort word_embeddings because load_from_dir may not load in alphabetic order. + # using a temporary copy so we don't reinitialize self.word_embeddings in case other objects have a reference to it. + sorted_word_embeddings = {e.name: e for e in sorted(self.word_embeddings.values(), key=lambda e: e.name.lower())} + self.word_embeddings.clear() + self.word_embeddings.update(sorted_word_embeddings) + + displayed_embeddings = (tuple(self.word_embeddings.keys()), tuple(self.skipped_embeddings.keys())) + if self.previously_displayed_embeddings != displayed_embeddings: + self.previously_displayed_embeddings = displayed_embeddings + t1 = time.time() + shared.log.info(f"Load embeddings: loaded={len(self.word_embeddings)} skipped={len(self.skipped_embeddings)} time={t1-t0:.2f}") + + + def find_embedding_at_position(self, tokens, offset): + token = tokens[offset] + possible_matches = self.ids_lookup.get(token, None) + if possible_matches is None: + return None, None + for ids, embedding in possible_matches: + if tokens[offset:offset + len(ids)] == ids: + return embedding, len(ids) + return None, None + + +def create_embedding(name, num_vectors_per_token, overwrite_old, init_text='*'): + cond_model = shared.sd_model.cond_stage_model + with devices.autocast(): + cond_model([""]) # will send cond model to GPU if lowvram/medvram is active + #cond_model expects at least some text, so we provide '*' as backup. + embedded = cond_model.encode_embedding_init_text(init_text or '*', num_vectors_per_token) + vec = torch.zeros((num_vectors_per_token, embedded.shape[1]), device=devices.device) + #Only copy if we provided an init_text, otherwise keep vectors as zeros + if init_text: + for i in range(num_vectors_per_token): + vec[i] = embedded[i * int(embedded.shape[0]) // num_vectors_per_token] + # Remove illegal characters from name. + name = "".join( x for x in name if (x.isalnum() or x in "._- ")) + fn = os.path.join(shared.opts.embeddings_dir, f"{name}.pt") + if not overwrite_old and os.path.exists(fn): + shared.log.warning(f"Embedding already exists: {fn}") + else: + embedding = Embedding(vec=vec, name=name, filename=fn) + embedding.step = 0 + embedding.save(fn) + shared.log.info(f'Created embedding: {fn} vectors {num_vectors_per_token} init {init_text}') + return fn + + +def write_loss(log_directory, filename, step, epoch_len, values): + if shared.opts.training_write_csv_every == 0: + return + if step % shared.opts.training_write_csv_every != 0: + return + write_csv_header = False if os.path.exists(os.path.join(log_directory, filename)) else True + with open(os.path.join(log_directory, filename), "a+", newline='', encoding='utf-8') as fout: + csv_writer = csv.DictWriter(fout, fieldnames=["step", "epoch", "epoch_step", *(values.keys())]) + if write_csv_header: + csv_writer.writeheader() + epoch = (step - 1) // epoch_len + epoch_step = (step - 1) % epoch_len + csv_writer.writerow({ + "step": step, + "epoch": epoch, + "epoch_step": epoch_step, + **values, + }) + + +def tensorboard_setup(log_directory): + os.makedirs(os.path.join(log_directory, "tensorboard"), exist_ok=True) + return SummaryWriter( + log_dir=os.path.join(log_directory, "tensorboard"), + flush_secs=shared.opts.training_tensorboard_flush_every) + + +def tensorboard_add(tensorboard_writer, loss, global_step, step, learn_rate, epoch_num): + tensorboard_add_scaler(tensorboard_writer, "Loss/train", loss, global_step) + tensorboard_add_scaler(tensorboard_writer, f"Loss/train/epoch-{epoch_num}", loss, step) + tensorboard_add_scaler(tensorboard_writer, "Learn rate/train", learn_rate, global_step) + tensorboard_add_scaler(tensorboard_writer, f"Learn rate/train/epoch-{epoch_num}", learn_rate, step) + + +def tensorboard_add_scaler(tensorboard_writer, tag, value, step): + tensorboard_writer.add_scalar(tag=tag, scalar_value=value, global_step=step) + + +def tensorboard_add_image(tensorboard_writer, tag, pil_image, step): + # Convert a pil image to a torch tensor + img_tensor = torch.as_tensor(np.array(pil_image, copy=True)) + img_tensor = img_tensor.view(pil_image.size[1], pil_image.size[0], len(pil_image.getbands())) + img_tensor = img_tensor.permute((2, 0, 1)) + tensorboard_writer.add_image(tag, img_tensor, global_step=step) + + +def validate_train_inputs(model_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_model_every, create_image_every, name="embedding"): + assert model_name, f"{name} not selected" + assert learn_rate, "Learning rate is empty or 0" + assert isinstance(batch_size, int), "Batch size must be integer" + assert batch_size > 0, "Batch size must be positive" + assert isinstance(gradient_step, int), "Gradient accumulation step must be integer" + assert gradient_step > 0, "Gradient accumulation step must be positive" + assert data_root, "Dataset directory is empty" + assert os.path.isdir(data_root), "Dataset directory doesn't exist" + assert os.listdir(data_root), "Dataset directory is empty" + assert template_filename, "Prompt template file not selected" + assert template_file, f"Prompt template file {template_filename} not found" + assert os.path.isfile(template_file.path), f"Prompt template file {template_filename} doesn't exist" + assert steps, "Max steps is empty or 0" + assert isinstance(steps, int), "Max steps must be integer" + assert steps > 0, "Max steps must be positive" + assert isinstance(save_model_every, int), "Save {name} must be integer" + assert save_model_every >= 0, "Save {name} must be positive or 0" + assert isinstance(create_image_every, int), "Create image must be integer" + assert create_image_every >= 0, "Create image must be positive or 0" + + +def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument + + shared.log.debug(f'train_embedding: embedding_name={embedding_name}|learn_rate={learn_rate}|batch_size={batch_size}|gradient_step={gradient_step}|data_root={data_root}|log_directory={log_directory}|training_width={training_width}|training_height={training_height}|varsize={varsize}|steps={steps}|clip_grad_mode={clip_grad_mode}|clip_grad_value={clip_grad_value}|shuffle_tags={shuffle_tags}|tag_drop_out={tag_drop_out}|latent_sampling_method={latent_sampling_method}|use_weight={use_weight}|create_image_every={create_image_every}|save_embedding_every={save_embedding_every}|template_filename={template_filename}|save_image_with_stored_embedding={save_image_with_stored_embedding}|preview_from_txt2img={preview_from_txt2img}|preview_prompt={preview_prompt}|preview_negative_prompt={preview_negative_prompt}|preview_steps={preview_steps}|preview_sampler_index={preview_sampler_index}|preview_cfg_scale={preview_cfg_scale}|preview_seed={preview_seed}|preview_width={preview_width}|preview_height={preview_height}') + save_embedding_every = save_embedding_every or 0 + create_image_every = create_image_every or 0 + template_file = textual_inversion_templates.get(template_filename, None) + validate_train_inputs(embedding_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_embedding_every, create_image_every, name="embedding") + if log_directory is None or log_directory == '': + log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}" + template_file = template_file.path + + shared.state.job = "train" + shared.state.textinfo = "Initializing textual inversion training..." + shared.state.job_count = steps + + filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt') + + if log_directory == '': + log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}" + log_directory = os.path.join(log_directory, embedding_name) + unload = shared.opts.unload_models_when_training + + if save_embedding_every > 0: + embedding_dir = os.path.join(log_directory, "embeddings") + os.makedirs(embedding_dir, exist_ok=True) + else: + embedding_dir = None + + if create_image_every > 0: + images_dir = os.path.join(log_directory, "images") + os.makedirs(images_dir, exist_ok=True) + else: + images_dir = None + + if create_image_every > 0 and save_image_with_stored_embedding: + images_embeds_dir = os.path.join(log_directory, "image_embeddings") + os.makedirs(images_embeds_dir, exist_ok=True) + else: + images_embeds_dir = None + + hijack = sd_hijack.model_hijack + embedding = hijack.embedding_db.word_embeddings[embedding_name] + checkpoint = sd_models.select_checkpoint() + initial_step = embedding.step or 0 + if initial_step >= steps: + shared.state.textinfo = "Model has already been trained beyond specified max steps" + return embedding, filename + scheduler = LearnRateScheduler(learn_rate, steps, initial_step) + clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else \ + torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else \ + None + if clip_grad: + clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False) + # dataset loading may take a while, so input validations and early returns should be done before this + shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..." + old_parallel_processing_allowed = shared.parallel_processing_allowed + + if shared.opts.training_enable_tensorboard: + tensorboard_writer = tensorboard_setup(log_directory) + + pin_memory = shared.opts.pin_memory + # init dataset + ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=embedding_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight) + + if shared.opts.save_training_settings_to_txt: + save_settings_to_file(log_directory, {**dict(model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds), num_vectors_per_token=len(embedding.vec)), **locals()}) + latent_sampling_method = ds.latent_sampling_method + # init dataloader + dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory) + if unload: + shared.parallel_processing_allowed = False + shared.sd_model.first_stage_model.to(devices.cpu) + + embedding.vec.requires_grad = True + optimizer = torch.optim.AdamW([embedding.vec], lr=scheduler.learn_rate, weight_decay=0.0) + if shared.opts.save_optimizer_state: + optimizer_state_dict = None + if os.path.exists(f"{filename}.optim"): + optimizer_saved_dict = torch.load(f"{filename}.optim", map_location='cpu') + if embedding.checksum() == optimizer_saved_dict.get('hash', None): + optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None) + if optimizer_state_dict is not None: + optimizer.load_state_dict(optimizer_state_dict) + shared.log.info("Load existing optimizer from checkpoint") + else: + shared.log.info("No saved optimizer exists in checkpoint") + + scaler = torch.cuda.amp.GradScaler() + + batch_size = ds.batch_size + gradient_step = ds.gradient_step + # n steps = batch_size * gradient_step * n image processed + steps_per_epoch = len(ds) // batch_size // gradient_step + max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step + loss_step = 0 + _loss_step = 0 #internal + last_saved_file = "" + last_saved_image = "" + forced_filename = "" + embedding_yet_to_be_embedded = False + is_training_inpainting_model = shared.sd_model.model.conditioning_key in {'hybrid', 'concat'} + img_c = None + + pbar = tqdm(total=steps - initial_step) + try: + sd_hijack_checkpoint.add() + for _i in range((steps-initial_step) * gradient_step): + if scheduler.finished: + break + if shared.state.interrupted: + break + for j, batch in enumerate(dl): + # works as a drop_last=True for gradient accumulation + if j == max_steps_per_epoch: + break + scheduler.apply(optimizer, embedding.step) + if scheduler.finished: + break + if shared.state.interrupted: + break + if clip_grad: + clip_grad_sched.step(embedding.step) + with devices.autocast(): + x = batch.latent_sample.to(devices.device, non_blocking=pin_memory) + if use_weight: + w = batch.weight.to(devices.device, non_blocking=pin_memory) + c = shared.sd_model.cond_stage_model(batch.cond_text) + if is_training_inpainting_model: + if img_c is None: + img_c = processing.txt2img_image_conditioning(shared.sd_model, c, training_width, training_height) + cond = {"c_concat": [img_c], "c_crossattn": [c]} + else: + cond = c + if use_weight: + loss = shared.sd_model.weighted_forward(x, cond, w)[0] / gradient_step + del w + else: + loss = shared.sd_model.forward(x, cond)[0] / gradient_step + del x + _loss_step += loss.item() + + scaler.scale(loss).backward() + # go back until we reach gradient accumulation steps + if (j + 1) % gradient_step != 0: + continue + if clip_grad: + clip_grad(embedding.vec, clip_grad_sched.learn_rate) + + scaler.step(optimizer) + scaler.update() + embedding.step += 1 + pbar.update() + optimizer.zero_grad(set_to_none=True) + loss_step = _loss_step + _loss_step = 0 + steps_done = embedding.step + 1 + epoch_num = embedding.step // steps_per_epoch + + description = f"Training textual inversion step {embedding.step} loss: {loss_step:.5f} lr: {scheduler.learn_rate:.5f}" + pbar.set_description(description) + if embedding_dir is not None and steps_done % save_embedding_every == 0: + # Before saving, change name to match current checkpoint. + embedding_name_every = f'{embedding_name}-{steps_done}' + last_saved_file = os.path.join(embedding_dir, f'{embedding_name_every}.pt') + save_embedding(embedding, optimizer, checkpoint, embedding_name_every, last_saved_file, remove_cached_checksum=True) + embedding_yet_to_be_embedded = True + + write_loss(log_directory, f"{embedding_name}.csv", embedding.step, steps_per_epoch, { "loss": f"{loss_step:.7f}", "learn_rate": scheduler.learn_rate }) + + if images_dir is not None and steps_done % create_image_every == 0: + forced_filename = f'{embedding_name}-{steps_done}' + last_saved_image = os.path.join(images_dir, forced_filename) + shared.sd_model.first_stage_model.to(devices.device) + + p = processing.StableDiffusionProcessingTxt2Img( + sd_model=shared.sd_model, + do_not_save_grid=True, + do_not_save_samples=True, + do_not_reload_embeddings=True, + ) + + if preview_from_txt2img: + p.prompt = preview_prompt + p.negative_prompt = preview_negative_prompt + p.steps = preview_steps + p.sampler_name = processing.get_sampler_name(preview_sampler_index) + p.cfg_scale = preview_cfg_scale + p.seed = preview_seed + p.width = preview_width + p.height = preview_height + else: + p.prompt = batch.cond_text[0] + p.steps = 20 + p.width = training_width + p.height = training_height + + preview_text = p.prompt + processed = processing.process_images(p) + image = processed.images[0] if len(processed.images) > 0 else None + + if unload: + shared.sd_model.first_stage_model.to(devices.cpu) + + if image is not None: + shared.state.assign_current_image(image) + last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) + last_saved_image += f", prompt: {preview_text}" + if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images: + tensorboard_add_image(tensorboard_writer, f"Validation at epoch {epoch_num}", image, embedding.step) + + if save_image_with_stored_embedding and os.path.exists(last_saved_file) and embedding_yet_to_be_embedded: + last_saved_image_chunks = os.path.join(images_embeds_dir, f'{embedding_name}-{steps_done}.png') + info = PngImagePlugin.PngInfo() + data = torch.load(last_saved_file) + info.add_text("sd-ti-embedding", embedding_to_b64(data)) + title = f"<{data.get('name', '???')}>" + try: + vectorSize = list(data['string_to_param'].values())[0].shape[0] + except Exception: + vectorSize = '?' + checkpoint = sd_models.select_checkpoint() + footer_left = checkpoint.model_name + footer_mid = f'[{checkpoint.shorthash}]' + footer_right = f'{vectorSize}v {steps_done}s' + captioned_image = caption_image_overlay(image, title, footer_left, footer_mid, footer_right) + captioned_image = insert_image_data_embed(captioned_image, data) + captioned_image.save(last_saved_image_chunks, "PNG", pnginfo=info) + embedding_yet_to_be_embedded = False + + last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) + last_saved_image += f", prompt: {preview_text}" + + shared.state.job_no = embedding.step + shared.state.textinfo = f""" +

+Loss: {loss_step:.7f}
+Step: {steps_done}
+Last prompt: {html.escape(batch.cond_text[0])}
+Last saved embedding: {html.escape(last_saved_file)}
+Last saved image: {html.escape(last_saved_image)}
+

+""" + filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt') + save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True) + except Exception as e: + errors.display(e, 'embedding train') + finally: + pbar.leave = False + pbar.close() + shared.sd_model.first_stage_model.to(devices.device) + shared.parallel_processing_allowed = old_parallel_processing_allowed + sd_hijack_checkpoint.remove() + return embedding, filename + + +def save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True): + old_embedding_name = embedding.name + old_sd_checkpoint = embedding.sd_checkpoint if hasattr(embedding, "sd_checkpoint") else None + old_sd_checkpoint_name = embedding.sd_checkpoint_name if hasattr(embedding, "sd_checkpoint_name") else None + old_cached_checksum = embedding.cached_checksum if hasattr(embedding, "cached_checksum") else None + try: + embedding.sd_checkpoint = checkpoint.shorthash + embedding.sd_checkpoint_name = checkpoint.model_name + if remove_cached_checksum: + embedding.cached_checksum = None + embedding.name = embedding_name + embedding.optimizer_state_dict = optimizer.state_dict() + embedding.save(filename) + except Exception: + embedding.sd_checkpoint = old_sd_checkpoint + embedding.sd_checkpoint_name = old_sd_checkpoint_name + embedding.name = old_embedding_name + embedding.cached_checksum = old_cached_checksum + raise diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 75dc94eb5..20907011a 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -1,826 +1,826 @@ -import io -import re -import time -import json -import html -import base64 -import os.path -import urllib.parse -import threading -from datetime import datetime -from types import SimpleNamespace -from pathlib import Path -from html.parser import HTMLParser -from collections import OrderedDict -import gradio as gr -from PIL import Image -from starlette.responses import FileResponse, JSONResponse -from modules import paths, shared, scripts, files_cache, errors -from modules.ui_components import ToolButton -import modules.ui_symbols as symbols - - -allowed_dirs = [] -dir_timestamps = {} -refresh_time = 0 -extra_pages = shared.extra_networks -debug = shared.log.trace if os.environ.get('SD_EN_DEBUG', None) is not None else lambda *args, **kwargs: None -debug('Trace: EN') -card_full = ''' -
-
-
-
{title}
-
-
- 🛈 -
    -
    - -
    -''' -card_list = ''' -
    -
    - 🛈  -
    {title}
      -
    -
    -
    -''' - - -def init_api(app): - - def fetch_file(filename: str = ""): - if not os.path.exists(filename): - return JSONResponse({ "error": f"file {filename}: not found" }, status_code=404) - if filename.startswith('html/') or filename.startswith('models/'): - return FileResponse(filename, headers={"Accept-Ranges": "bytes"}) - if not any(Path(folder).absolute() in Path(filename).absolute().parents for folder in allowed_dirs): - return JSONResponse({ "error": f"file {filename}: must be in one of allowed directories" }, status_code=403) - if os.path.splitext(filename)[1].lower() not in (".png", ".jpg", ".jpeg", ".webp"): - return JSONResponse({"error": f"file {filename}: not an image file"}, status_code=403) - return FileResponse(filename, headers={"Accept-Ranges": "bytes"}) - - def get_metadata(page: str = "", item: str = ""): - page = next(iter([x for x in shared.extra_networks if x.name == page]), None) - if page is None: - return JSONResponse({ 'metadata': 'none' }) - metadata = page.metadata.get(item, 'none') - if metadata is None: - metadata = '' - # shared.log.debug(f"Extra networks metadata: page='{page}' item={item} len={len(metadata)}") - return JSONResponse({"metadata": metadata}) - - def get_info(page: str = "", item: str = ""): - page = next(iter([x for x in get_pages() if x.name == page]), None) - if page is None: - return JSONResponse({ 'info': 'none' }) - item = next(iter([x for x in page.items if x['name'] == item]), None) - if item is None: - return JSONResponse({ 'info': 'none' }) - info = page.find_info(item['filename']) - if info is None: - info = {} - # shared.log.debug(f"Extra networks info: page='{page.name}' item={item['name']} len={len(info)}") - return JSONResponse({"info": info}) - - def get_desc(page: str = "", item: str = ""): - page = next(iter([x for x in get_pages() if x.name == page]), None) - if page is None: - return JSONResponse({ 'description': 'none' }) - item = next(iter([x for x in page.items if x['name'] == item]), None) - if item is None: - return JSONResponse({ 'description': 'none' }) - desc = page.find_description(item['filename']) - if desc is None: - desc = '' - # shared.log.debug(f"Extra networks desc: page='{page.name}' item={item['name']} len={len(desc)}") - return JSONResponse({"description": desc}) - - app.add_api_route("/sd_extra_networks/thumb", fetch_file, methods=["GET"]) - app.add_api_route("/sd_extra_networks/metadata", get_metadata, methods=["GET"]) - app.add_api_route("/sd_extra_networks/info", get_info, methods=["GET"]) - app.add_api_route("/sd_extra_networks/description", get_desc, methods=["GET"]) - - -class ExtraNetworksPage: - def __init__(self, title): - self.title = title - self.name = title.lower() - self.allow_negative_prompt = False - self.metadata = {} - self.info = {} - self.html = '' - self.items = [] - self.missing_thumbs = [] - self.refresh_time = 0 - self.page_time = 0 - self.list_time = 0 - self.info_time = 0 - self.desc_time = 0 - self.dirs = {} - self.view = shared.opts.extra_networks_view - self.card = card_full if shared.opts.extra_networks_view == 'gallery' else card_list - - def refresh(self): - pass - - def create_xyz_grid(self): - xyz_grid = [x for x in scripts.scripts_data if x.script_class.__module__ == "xyz_grid.py"][0].module - - def add_prompt(p, opt, x): - for item in [x for x in self.items if x["name"] == opt]: - try: - p.prompt = f'{p.prompt} {eval(item["prompt"])}' # pylint: disable=eval-used - except Exception as e: - shared.log.error(f'Cannot evaluate extra network prompt: {item["prompt"]} {e}') - - if not any(self.title in x.label for x in xyz_grid.axis_options): - if self.title == 'Model': - return - opt = xyz_grid.AxisOption(f"[Network] {self.title}", str, add_prompt, choices=lambda: [x["name"] for x in self.items]) - xyz_grid.axis_options.append(opt) - - def link_preview(self, filename): - quoted_filename = urllib.parse.quote(filename.replace('\\', '/')) - mtime = os.path.getmtime(filename) - preview = f"./sd_extra_networks/thumb?filename={quoted_filename}&mtime={mtime}" - return preview - - def search_terms_from_path(self, filename): - return filename.replace('\\', '/') - - def is_empty(self, folder): - return any(files_cache.list_files(folder, ext_filter=['.ckpt', '.safetensors', '.pt', '.json'])) - - def create_thumb(self): - debug(f'EN create-thumb: {self.name}') - created = 0 - for f in self.missing_thumbs: - if not os.path.exists(f): - continue - fn, _ext = os.path.splitext(f) - fn = fn.replace('.preview', '') - fn = f'{fn}.thumb.jpg' - if os.path.exists(fn): - continue - img = None - try: - img = Image.open(f) - except Exception: - img = None - shared.log.warning(f'Extra network removing invalid image: {f}') - try: - if img is None: - img = None - os.remove(f) - elif img.width > 1024 or img.height > 1024 or os.path.getsize(f) > 65536: - img = img.convert('RGB') - img.thumbnail((512, 512), Image.Resampling.HAMMING) - img.save(fn, quality=50) - img.close() - created += 1 - except Exception as e: - shared.log.warning(f'Extra network error creating thumbnail: {f} {e}') - if created > 0: - shared.log.info(f"Extra network thumbnails: {self.name} created={created}") - self.missing_thumbs.clear() - - def create_items(self, tabname): - if self.refresh_time is not None and self.refresh_time > refresh_time: # cached results - return - t0 = time.time() - try: - self.items = list(self.list_items()) - self.refresh_time = time.time() - except Exception as e: - self.items = [] - shared.log.error(f'Extra networks error listing items: class={self.__class__.__name__} tab={tabname} {e}') - for item in self.items: - if item is None: - continue - self.metadata[item["name"]] = item.get("metadata", {}) - t1 = time.time() - debug(f'EN create-items: page={self.name} items={len(self.items)} time={t1-t0:.2f}') - self.list_time += t1-t0 - - - def create_page(self, tabname, skip = False): - debug(f'EN create-page: {self.name}') - if self.page_time > refresh_time and len(self.html) > 0: # cached page - return self.html - self_name_id = self.name.replace(" ", "_") - if skip: - return f"
    Extra network page not ready
    Click refresh to try again
    " - subdirs = {} - allowed_folders = [os.path.abspath(x) for x in self.allowed_directories_for_previews()] - for parentdir, dirs in {d: files_cache.walk(d, cached=True, recurse=files_cache.not_hidden) for d in allowed_folders}.items(): - for tgt in dirs: - tgt = tgt.path - if os.path.join(paths.models_path, 'Reference') in tgt: - subdirs['Reference'] = 1 - if shared.backend == shared.Backend.DIFFUSERS and shared.opts.diffusers_dir in tgt: - subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 - if 'models--' in tgt: - continue - subdir = tgt[len(parentdir):].replace("\\", "/") - while subdir.startswith("/"): - subdir = subdir[1:] - if not subdir: - continue - # if not self.is_empty(tgt): - subdirs[subdir] = 1 - debug(f"Extra networks: page='{self.name}' subfolders={list(subdirs)}") - subdirs = OrderedDict(sorted(subdirs.items())) - if self.name == 'model': - subdirs['Reference'] = 1 - subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 - subdirs.move_to_end(os.path.basename(shared.opts.diffusers_dir)) - subdirs.move_to_end('Reference') - if self.name == 'style' and shared.opts.extra_networks_styles: - subdirs['built-in'] = 1 - subdirs_html = "
    " - subdirs_html += "".join([f"
    " for subdir in subdirs if subdir != '']) - self.html = '' - self.create_items(tabname) - self.create_xyz_grid() - htmls = [] - if len(self.items) > 0 and self.items[0].get('mtime', None) is not None: - self.items.sort(key=lambda x: x["mtime"], reverse=True) - for item in self.items: - htmls.append(self.create_html(item, tabname)) - self.html += ''.join(htmls) - self.page_time = time.time() - if len(subdirs_html) > 0 or len(self.html) > 0: - self.html = f"
    {subdirs_html}
    {self.html}
    " - else: - return '' - shared.log.debug(f"Extra networks: page='{self.name}' items={len(self.items)} subfolders={len(subdirs)} tab={tabname} folders={self.allowed_directories_for_previews()} list={self.list_time:.2f} desc={self.desc_time:.2f} info={self.info_time:.2f} workers={shared.max_workers}") - if len(self.missing_thumbs) > 0: - threading.Thread(target=self.create_thumb).start() - return self.html - - def list_items(self): - raise NotImplementedError - - def allowed_directories_for_previews(self): - return [] - - def create_html(self, item, tabname): - try: - args = { - "tabname": tabname, - "page": self.name, - "name": item["name"], - "title": os.path.basename(item["name"].replace('_', ' ')), - "filename": item["filename"], - "tags": '|'.join([item.get("tags")] if isinstance(item.get("tags", {}), str) else list(item.get("tags", {}).keys())), - "preview": html.escape(item.get("preview", self.link_preview('html/card-no-preview.png'))), - "width": shared.opts.extra_networks_card_size, - "height": shared.opts.extra_networks_card_size if shared.opts.extra_networks_card_square else 'auto', - "fit": shared.opts.extra_networks_card_fit, - "prompt": item.get("prompt", None), - "search": item.get("search_term", ""), - "description": item.get("description") or "", - "card_click": item.get("onclick", '"' + html.escape(f'return cardClicked({item.get("prompt", None)}, {"true" if self.allow_negative_prompt else "false"})') + '"'), - "mtime": item.get("mtime", 0), - "size": item.get("size", 0), - } - alias = item.get("alias", None) - if alias is not None: - args['title'] += f'\nAlias: {alias}' - return self.card.format(**args) - except Exception as e: - shared.log.error(f'Extra networks item error: page={tabname} item={item["name"]} {e}') - return "" - - def find_preview_file(self, path): - exts = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] - if path is None: - return 'html/card-no-preview.png' - if shared.opts.diffusers_dir in path: - path = os.path.relpath(path, shared.opts.diffusers_dir) - ref = os.path.join('models', 'Reference') - fn = os.path.join(ref, path.replace('models--', '').replace('\\', '/').split('/')[0]) - files = list(files_cache.list_files(ref, ext_filter=exts, recursive=False)) - else: - files = list(files_cache.list_files(os.path.dirname(path), ext_filter=exts, recursive=False)) - fn = os.path.splitext(path)[0] - for file in [f'{fn}{mid}{ext}' for ext in exts for mid in ['.thumb.', '.', '.preview.']]: - if file in files: - if 'Reference' not in file and '.thumb.' not in file: - self.missing_thumbs.append(file) - return file - return 'html/card-no-preview.png' - - def find_preview(self, path): - preview_file = self.find_preview_file(path) - return self.link_preview(preview_file) - - def find_description(self, path, info=None): - t0 = time.time() - class HTMLFilter(HTMLParser): - text = "" - def handle_data(self, data): - self.text += data - def handle_endtag(self, tag): - if tag == 'p': - self.text += '\n' - - fn = os.path.splitext(path)[0] + '.txt' - if os.path.exists(fn): - try: - with open(fn, "r", encoding="utf-8", errors="replace") as f: - txt = f.read() - txt = re.sub('[<>]', '', txt) - return txt - except OSError: - pass - if info is None: - info = self.find_info(path) - desc = info.get('description', '') or '' - f = HTMLFilter() - f.feed(desc) - t1 = time.time() - self.desc_time += t1-t0 - return f.text - - def find_info(self, path): - fn = os.path.splitext(path)[0] + '.json' - data = {} - if os.path.exists(fn): - t0 = time.time() - data = shared.readfile(fn, silent=True) - if type(data) is list: - data = data[0] - t1 = time.time() - self.info_time += t1-t0 - return data - - -def initialize(): - shared.extra_networks.clear() - - -def register_page(page: ExtraNetworksPage): - # registers extra networks page for the UI; recommend doing it in on_before_ui() callback for extensions - debug(f'EN register-page: {page}') - if page in shared.extra_networks: - debug(f'EN register-page: {page} already registered') - return - shared.extra_networks.append(page) - # allowed_dirs.clear() - # for pg in shared.extra_networks: - for folder in page.allowed_directories_for_previews(): - if folder not in allowed_dirs: - allowed_dirs.append(os.path.abspath(folder)) - - -def register_pages(): - from modules.ui_extra_networks_textual_inversion import ExtraNetworksPageTextualInversion - from modules.ui_extra_networks_hypernets import ExtraNetworksPageHypernetworks - from modules.ui_extra_networks_checkpoints import ExtraNetworksPageCheckpoints - from modules.ui_extra_networks_styles import ExtraNetworksPageStyles - from modules.ui_extra_networks_vae import ExtraNetworksPageVAEs - debug('EN register-pages') - register_page(ExtraNetworksPageCheckpoints()) - register_page(ExtraNetworksPageStyles()) - register_page(ExtraNetworksPageTextualInversion()) - register_page(ExtraNetworksPageHypernetworks()) - register_page(ExtraNetworksPageVAEs()) - - -def get_pages(title=None): - pages = [] - if 'All' in shared.opts.extra_networks: - pages = shared.extra_networks - else: - titles = [page.title for page in shared.extra_networks] - if title is None: - for page in shared.opts.extra_networks: - try: - idx = titles.index(page) - pages.append(shared.extra_networks[idx]) - except ValueError: - continue - else: - try: - idx = titles.index(title) - pages.append(shared.extra_networks[idx]) - except ValueError: - pass - return pages - - -class ExtraNetworksUi: - def __init__(self): - self.tabname: str = None - self.pages: list(str) = None - self.visible: gr.State = None - self.state: gr.Textbox = None - self.details: gr.Group = None - self.tabs: gr.Tabs = None - self.gallery: gr.Gallery = None - self.description: gr.Textbox = None - self.search: gr.Textbox = None - self.button_details: gr.Button = None - self.button_refresh: gr.Button = None - self.button_scan: gr.Button = None - self.button_view: gr.Button = None - self.button_quicksave: gr.Button = None - self.button_save: gr.Button = None - self.button_sort: gr.Button = None - self.button_apply: gr.Button = None - self.button_close: gr.Button = None - self.button_model: gr.Checkbox = None - self.details_components: list = [] - self.last_item: dict = None - self.last_page: ExtraNetworksPage = None - self.state: gr.State = None - - -def create_ui(container, button_parent, tabname, skip_indexing = False): - debug(f'EN create-ui: {tabname}') - ui = ExtraNetworksUi() - ui.tabname = tabname - ui.pages = [] - ui.state = gr.Textbox('{}', elem_id=f"{tabname}_extra_state", visible=False) - ui.visible = gr.State(value=False) # pylint: disable=abstract-class-instantiated - ui.details = gr.Group(elem_id=f"{tabname}_extra_details", visible=False) - ui.tabs = gr.Tabs(elem_id=f"{tabname}_extra_tabs") - ui.button_details = gr.Button('Details', elem_id=f"{tabname}_extra_details_btn", visible=False) - state = {} - if shared.cmd_opts.profile: - import cProfile - pr = cProfile.Profile() - pr.enable() - - def get_item(state, params = None): - if params is not None and type(params) == dict: - page = next(iter([x for x in get_pages() if x.title == 'Style']), None) - item = page.create_style(params) - else: - if state is None or not hasattr(state, 'page') or not hasattr(state, 'item'): - return None, None - page = next(iter([x for x in get_pages() if x.title == state.page]), None) - if page is None: - return None, None - item = next(iter([x for x in page.items if x["name"] == state.item]), None) - if item is None: - return page, None - item = SimpleNamespace(**item) - ui.last_item = item - ui.last_page = page - return page, item - - # main event that is triggered when js updates state text field with json values, used to communicate js -> python - def state_change(state_text): - try: - nonlocal state - state = SimpleNamespace(**json.loads(state_text)) - except Exception as e: - shared.log.error(f'Extra networks state error: {e}') - return - _page, _item = get_item(state) - # shared.log.debug(f'Extra network: op={state.op} page={page.title if page is not None else None} item={item.filename if item is not None else None}') - - def toggle_visibility(is_visible): - is_visible = not is_visible - return is_visible, gr.update(visible=is_visible), gr.update(variant=("secondary-down" if is_visible else "secondary")) - - with ui.details: - details_close = ToolButton(symbols.close, elem_id=f"{tabname}_extra_details_close", elem_classes=['extra-details-close']) - details_close.click(fn=lambda: gr.update(visible=False), inputs=[], outputs=[ui.details]) - with gr.Row(): - with gr.Column(scale=1): - text = gr.HTML('
    title
    ') - ui.details_components.append(text) - with gr.Column(scale=1): - img = gr.Image(value=None, show_label=False, interactive=False, container=False, show_download_button=False, show_info=False, elem_id=f"{tabname}_extra_details_img", elem_classes=['extra-details-img']) - ui.details_components.append(img) - with gr.Row(): - btn_save_img = gr.Button('Replace', elem_classes=['small-button']) - btn_delete_img = gr.Button('Delete', elem_classes=['small-button']) - with gr.Tabs(): - with gr.Tab('Description'): - desc = gr.Textbox('', show_label=False, lines=8, placeholder="Extra network description...") - ui.details_components.append(desc) - with gr.Row(): - btn_save_desc = gr.Button('Save', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_save_desc') - btn_delete_desc = gr.Button('Delete', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_delete_desc') - btn_close_desc = gr.Button('Close', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_close_desc') - btn_close_desc.click(fn=lambda: gr.update(visible=False), _js='refeshDetailsEN', inputs=[], outputs=[ui.details]) - with gr.Tab('Model metadata'): - info = gr.JSON({}, show_label=False) - ui.details_components.append(info) - with gr.Row(): - btn_save_info = gr.Button('Save', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_save_info') - btn_delete_info = gr.Button('Delete', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_delete_info') - btn_close_info = gr.Button('Close', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_close_info') - btn_close_info.click(fn=lambda: gr.update(visible=False), _js='refeshDetailsEN', inputs=[], outputs=[ui.details]) - with gr.Tab('Embedded metadata'): - meta = gr.JSON({}, show_label=False) - ui.details_components.append(meta) - - with ui.tabs: - def ui_tab_change(page): - scan_visible = page in ['Model', 'Lora', 'Hypernetwork', 'Embedding'] - save_visible = page in ['Style'] - model_visible = page in ['Model'] - return [gr.update(visible=scan_visible), gr.update(visible=save_visible), gr.update(visible=model_visible)] - - ui.button_refresh = ToolButton(symbols.refresh, elem_id=f"{tabname}_extra_refresh") - ui.button_scan = ToolButton(symbols.scan, elem_id=f"{tabname}_extra_scan", visible=True) - ui.button_quicksave = ToolButton(symbols.book, elem_id=f"{tabname}_extra_quicksave", visible=False) - ui.button_save = ToolButton(symbols.book, elem_id=f"{tabname}_extra_save", visible=False) - ui.button_sort = ToolButton(symbols.sort, elem_id=f"{tabname}_extra_sort", visible=True) - ui.button_view = ToolButton(symbols.view, elem_id=f"{tabname}_extra_view", visible=True) - ui.button_close = ToolButton(symbols.close, elem_id=f"{tabname}_extra_close", visible=True) - ui.button_model = ToolButton(symbols.refine, elem_id=f"{tabname}_extra_model", visible=True) - ui.search = gr.Textbox('', show_label=False, elem_id=f"{tabname}_extra_search", placeholder="Search...", elem_classes="textbox", lines=2, container=False) - ui.description = gr.Textbox('', show_label=False, elem_id=f"{tabname}_description", elem_classes="textbox", lines=2, interactive=False, container=False) - - if ui.tabname == 'txt2img': # refresh only once - global refresh_time # pylint: disable=global-statement - refresh_time = time.time() - if not skip_indexing: - threads = [] - for page in get_pages(): - if os.environ.get('SD_EN_DEBUG', None) is not None: - threads.append(threading.Thread(target=page.create_items, args=[ui.tabname])) - threads[-1].start() - else: - page.create_items(ui.tabname) - for thread in threads: - thread.join() - for page in get_pages(): - page.create_page(ui.tabname, skip_indexing) - with gr.Tab(page.title, id=page.title.lower().replace(" ", "_"), elem_classes="extra-networks-tab") as tab: - page_html = gr.HTML(page.html, elem_id=f'{tabname}{page.name}_extra_page', elem_classes="extra-networks-page") - ui.pages.append(page_html) - tab.select(ui_tab_change, _js="getENActivePage", inputs=[ui.button_details], outputs=[ui.button_scan, ui.button_save, ui.button_model]) - if shared.cmd_opts.profile: - errors.profile(pr, 'ExtraNetworks') - pr.disable() - # ui.tabs.change(fn=ui_tab_change, inputs=[], outputs=[ui.button_scan, ui.button_save]) - - def fn_save_img(image): - if ui.last_item is None or ui.last_item.local_preview is None: - return 'html/card-no-preview.png' - images = list(ui.gallery.temp_files) # gallery cannot be used as input component so looking at most recently registered temp files - if len(images) < 1: - shared.log.warning(f'Extra network no image: item={ui.last_item.name}') - return 'html/card-no-preview.png' - try: - images.sort(key=lambda f: os.path.getmtime(f), reverse=True) - image = Image.open(images[0]) - except Exception as e: - shared.log.error(f'Extra network error opening image: item={ui.last_item.name} {e}') - return 'html/card-no-preview.png' - fn_delete_img(image) - if image.width > 512 or image.height > 512: - image = image.convert('RGB') - image.thumbnail((512, 512), Image.Resampling.HAMMING) - try: - image.save(ui.last_item.local_preview, quality=50) - shared.log.debug(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}"') - except Exception as e: - shared.log.error(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}" {e}') - return image - - def fn_delete_img(_image): - preview_extensions = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] - fn = os.path.splitext(ui.last_item.filename)[0] - for file in [f'{fn}{mid}{ext}' for ext in preview_extensions for mid in ['.thumb.', '.preview.', '.']]: - if os.path.exists(file): - os.remove(file) - shared.log.debug(f'Extra network delete image: item={ui.last_item.name} filename="{file}"') - return 'html/card-no-preview.png' - - def fn_save_desc(desc): - if hasattr(ui.last_item, 'type') and ui.last_item.type == 'Style': - params = ui.last_page.parse_desc(desc) - if params is not None: - fn_save_info(params) - else: - fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' - with open(fn, 'w', encoding='utf-8') as f: - f.write(desc) - shared.log.debug(f'Extra network save desc: item={ui.last_item.name} filename="{fn}"') - return desc - - def fn_delete_desc(desc): - if ui.last_item is None: - return desc - if hasattr(ui.last_item, 'type') and ui.last_item.type == 'Style': - fn = os.path.splitext(ui.last_item.filename)[0] + '.json' - else: - fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' - if os.path.exists(fn): - shared.log.debug(f'Extra network delete desc: item={ui.last_item.name} filename="{fn}"') - os.remove(fn) - return '' - return desc - - def fn_save_info(info): - fn = os.path.splitext(ui.last_item.filename)[0] + '.json' - shared.writefile(info, fn, silent=True) - shared.log.debug(f'Extra network save info: item={ui.last_item.name} filename="{fn}"') - return info - - def fn_delete_info(info): - if ui.last_item is None: - return info - fn = os.path.splitext(ui.last_item.filename)[0] + '.json' - if os.path.exists(fn): - shared.log.debug(f'Extra network delete info: item={ui.last_item.name} filename="{fn}"') - os.remove(fn) - return '' - return info - - btn_save_img.click(fn=fn_save_img, _js='closeDetailsEN', inputs=[img], outputs=[img]) - btn_delete_img.click(fn=fn_delete_img, _js='closeDetailsEN', inputs=[img], outputs=[img]) - btn_save_desc.click(fn=fn_save_desc, _js='closeDetailsEN', inputs=[desc], outputs=[desc]) - btn_delete_desc.click(fn=fn_delete_desc, _js='closeDetailsEN', inputs=[desc], outputs=[desc]) - btn_save_info.click(fn=fn_save_info, _js='closeDetailsEN', inputs=[info], outputs=[info]) - btn_delete_info.click(fn=fn_delete_info, _js='closeDetailsEN', inputs=[info], outputs=[info]) - - def show_details(text, img, desc, info, meta, params): - page, item = get_item(state, params) - if item is not None and hasattr(item, 'name'): - stat = os.stat(item.filename) if os.path.exists(item.filename) else None - desc = item.description - fullinfo = shared.readfile(os.path.splitext(item.filename)[0] + '.json', silent=True) - if 'modelVersions' in fullinfo: # sanitize massive objects - fullinfo['modelVersions'] = [] - info = fullinfo - meta = page.metadata.get(item.name, {}) or {} - if type(meta) is str: - try: - meta = json.loads(meta) - except Exception: - meta = {} - if ui.last_item.preview.startswith('data:'): - b64str = ui.last_item.preview.split(',',1)[1] - img = Image.open(io.BytesIO(base64.b64decode(b64str))) - elif hasattr(item, 'local_preview') and os.path.exists(item.local_preview): - img = item.local_preview - else: - img = page.find_preview_file(item.filename) - lora = '' - model = '' - style = '' - note = '' - if not os.path.exists(item.filename): - note = f'
    Target filename: {item.filename}' - if page.title == 'Model': - merge = len(list(meta.get('sd_merge_models', {}))) - if merge > 0: - model += f'Merge models{merge} recipes' - if meta.get('modelspec.architecture', None) is not None: - model += f''' - Architecture{meta.get('modelspec.architecture', 'N/A')} - Title{meta.get('modelspec.title', 'N/A')} - Resolution{meta.get('modelspec.resolution', 'N/A')} - ''' - if page.title == 'Lora': - try: - tags = getattr(item, 'tags', {}) - tags = [f'{name}:{tags[name]}' for i, name in enumerate(tags)] - tags = ' '.join(tags) - except Exception: - tags = '' - try: - triggers = ' '.join(info.get('tags', [])) - except Exception: - triggers = '' - lora = f''' - Model tags{tags} - User tags{triggers} - Base model{meta.get('ss_sd_model_name', 'N/A')} - Resolution{meta.get('ss_resolution', 'N/A')} - Training images{meta.get('ss_num_train_images', 'N/A')} - Comment{meta.get('ss_training_comment', 'N/A')} - ''' - if page.title == 'Style': - style = f''' - Name{item.name} - Description{item.description} - Preview Embedded{item.preview.startswith('data:')} - ''' - desc = f'Name: {os.path.basename(item.name)}\nDescription: {item.description}\nPrompt: {item.prompt}\nNegative: {item.negative}\nExtra: {item.extra}\n' - text = f''' -

    {item.name}

    - - - - - - - - - {lora} - {model} - {style} -
    Type{page.title}
    Alias{getattr(item, 'alias', 'N/A')}
    Filename{item.filename}
    Hash{getattr(item, 'hash', 'N/A')}
    Size{round(stat.st_size/1024/1024, 2) if stat is not None else 'N/A'} MB
    Last modified{datetime.fromtimestamp(stat.st_mtime) if stat is not None else 'N/A'}
    - {note} - ''' - return [text, img, desc, info, meta, gr.update(visible=item is not None)] - - def ui_refresh_click(title): - pages = [] - for page in get_pages(): - if title is None or title == '' or title == page.title or len(page.html) == 0: - page.page_time = 0 - page.refresh_time = 0 - page.refresh() - page.create_page(ui.tabname) - shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname}") - pages.append(page.html) - ui.search.update(value = ui.search.value) - return pages - - def ui_view_cards(title): - pages = [] - for page in get_pages(): - if title is None or title == '' or title == page.title or len(page.html) == 0: - shared.opts.extra_networks_view = page.view - page.view = 'gallery' if page.view == 'list' else 'list' - page.card = card_full if page.view == 'gallery' else card_list - page.html = '' - page.create_page(ui.tabname) - shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname} view={page.view}") - pages.append(page.html) - ui.search.update(value = ui.search.value) - return pages - - def ui_scan_click(title): - from modules import ui_models - if ui_models.search_metadata_civit is not None: - ui_models.search_metadata_civit(True, title) - return ui_refresh_click(title) - - def ui_save_click(): - from modules import generation_parameters_copypaste - filename = os.path.join(paths.data_path, "params.txt") - if os.path.exists(filename): - with open(filename, "r", encoding="utf8") as file: - prompt = file.read() - else: - prompt = '' - params = generation_parameters_copypaste.parse_generation_parameters(prompt) - res = show_details(text=None, img=None, desc=None, info=None, meta=None, params=params) - return res - - def ui_quicksave_click(name): - from modules import generation_parameters_copypaste - fn = os.path.join(paths.data_path, "params.txt") - if os.path.exists(fn): - with open(fn, "r", encoding="utf8") as file: - prompt = file.read() - else: - prompt = '' - params = generation_parameters_copypaste.parse_generation_parameters(prompt) - fn = os.path.join(shared.opts.styles_dir, os.path.splitext(name)[0] + '.json') - prompt = params.get('Prompt', '') - item = { - "name": name, - "description": '', - "prompt": prompt, - "negative": params.get('Negative prompt', ''), - "extra": '', - # "type": 'Style', - # "title": name, - # "filename": fn, - # "search_term": None, - # "preview": None, - # "local_preview": None, - } - shared.writefile(item, fn, silent=True) - if len(prompt) > 0: - shared.log.debug(f"Extra network quick save style: item={name} filename='{fn}'") - else: - shared.log.warning(f"Extra network quick save model: item={name} filename='{fn}' prompt is empty") - - def ui_sort_cards(msg): - shared.log.debug(f'Extra networks: {msg}') - return msg - - dummy = gr.State(value=False) # pylint: disable=abstract-class-instantiated - button_parent.click(fn=toggle_visibility, inputs=[ui.visible], outputs=[ui.visible, container, button_parent]) - ui.button_close.click(fn=toggle_visibility, inputs=[ui.visible], outputs=[ui.visible, container]) - ui.button_sort.click(fn=ui_sort_cards, _js='sortExtraNetworks', inputs=[ui.search], outputs=[ui.description]) - ui.button_view.click(fn=ui_view_cards, inputs=[ui.search], outputs=ui.pages) - ui.button_refresh.click(fn=ui_refresh_click, _js='getENActivePage', inputs=[ui.search], outputs=ui.pages) - ui.button_scan.click(fn=ui_scan_click, _js='getENActivePage', inputs=[ui.search], outputs=ui.pages) - ui.button_save.click(fn=ui_save_click, inputs=[], outputs=ui.details_components + [ui.details]) - ui.button_quicksave.click(fn=ui_quicksave_click, _js="() => prompt('Prompt name', '')", inputs=[ui.search], outputs=[]) - ui.button_details.click(show_details, _js="getCardDetails", inputs=ui.details_components + [dummy], outputs=ui.details_components + [ui.details]) - ui.state.change(state_change, inputs=[ui.state], outputs=[]) - return ui - - -def setup_ui(ui, gallery): - ui.gallery = gallery +import io +import re +import time +import json +import html +import base64 +import os.path +import urllib.parse +import threading +from datetime import datetime +from types import SimpleNamespace +from pathlib import Path +from html.parser import HTMLParser +from collections import OrderedDict +import gradio as gr +from PIL import Image +from starlette.responses import FileResponse, JSONResponse +from modules import paths, shared, scripts, files_cache, errors +from modules.ui_components import ToolButton +import modules.ui_symbols as symbols + + +allowed_dirs = [] +dir_timestamps = {} +refresh_time = 0 +extra_pages = shared.extra_networks +debug = shared.log.trace if os.environ.get('SD_EN_DEBUG', None) is not None else lambda *args, **kwargs: None +debug('Trace: EN') +card_full = ''' +
    +
    +
    +
    {title}
    +
    +
    + 🛈 +
      +
      + +
      +''' +card_list = ''' +
      +
      + 🛈  +
      {title}
        +
      +
      +
      +''' + + +def init_api(app): + + def fetch_file(filename: str = ""): + if not os.path.exists(filename): + return JSONResponse({ "error": f"file {filename}: not found" }, status_code=404) + if filename.startswith('html/') or filename.startswith('models/'): + return FileResponse(filename, headers={"Accept-Ranges": "bytes"}) + if not any(Path(folder).absolute() in Path(filename).absolute().parents for folder in allowed_dirs): + return JSONResponse({ "error": f"file {filename}: must be in one of allowed directories" }, status_code=403) + if os.path.splitext(filename)[1].lower() not in (".png", ".jpg", ".jpeg", ".webp"): + return JSONResponse({"error": f"file {filename}: not an image file"}, status_code=403) + return FileResponse(filename, headers={"Accept-Ranges": "bytes"}) + + def get_metadata(page: str = "", item: str = ""): + page = next(iter([x for x in shared.extra_networks if x.name == page]), None) + if page is None: + return JSONResponse({ 'metadata': 'none' }) + metadata = page.metadata.get(item, 'none') + if metadata is None: + metadata = '' + # shared.log.debug(f"Extra networks metadata: page='{page}' item={item} len={len(metadata)}") + return JSONResponse({"metadata": metadata}) + + def get_info(page: str = "", item: str = ""): + page = next(iter([x for x in get_pages() if x.name == page]), None) + if page is None: + return JSONResponse({ 'info': 'none' }) + item = next(iter([x for x in page.items if x['name'] == item]), None) + if item is None: + return JSONResponse({ 'info': 'none' }) + info = page.find_info(item['filename']) + if info is None: + info = {} + # shared.log.debug(f"Extra networks info: page='{page.name}' item={item['name']} len={len(info)}") + return JSONResponse({"info": info}) + + def get_desc(page: str = "", item: str = ""): + page = next(iter([x for x in get_pages() if x.name == page]), None) + if page is None: + return JSONResponse({ 'description': 'none' }) + item = next(iter([x for x in page.items if x['name'] == item]), None) + if item is None: + return JSONResponse({ 'description': 'none' }) + desc = page.find_description(item['filename']) + if desc is None: + desc = '' + # shared.log.debug(f"Extra networks desc: page='{page.name}' item={item['name']} len={len(desc)}") + return JSONResponse({"description": desc}) + + app.add_api_route("/sd_extra_networks/thumb", fetch_file, methods=["GET"]) + app.add_api_route("/sd_extra_networks/metadata", get_metadata, methods=["GET"]) + app.add_api_route("/sd_extra_networks/info", get_info, methods=["GET"]) + app.add_api_route("/sd_extra_networks/description", get_desc, methods=["GET"]) + + +class ExtraNetworksPage: + def __init__(self, title): + self.title = title + self.name = title.lower() + self.allow_negative_prompt = False + self.metadata = {} + self.info = {} + self.html = '' + self.items = [] + self.missing_thumbs = [] + self.refresh_time = 0 + self.page_time = 0 + self.list_time = 0 + self.info_time = 0 + self.desc_time = 0 + self.dirs = {} + self.view = shared.opts.extra_networks_view + self.card = card_full if shared.opts.extra_networks_view == 'gallery' else card_list + + def refresh(self): + pass + + def create_xyz_grid(self): + xyz_grid = [x for x in scripts.scripts_data if x.script_class.__module__ == "xyz_grid.py"][0].module + + def add_prompt(p, opt, x): + for item in [x for x in self.items if x["name"] == opt]: + try: + p.prompt = f'{p.prompt} {eval(item["prompt"])}' # pylint: disable=eval-used + except Exception as e: + shared.log.error(f'Cannot evaluate extra network prompt: {item["prompt"]} {e}') + + if not any(self.title in x.label for x in xyz_grid.axis_options): + if self.title == 'Model': + return + opt = xyz_grid.AxisOption(f"[Network] {self.title}", str, add_prompt, choices=lambda: [x["name"] for x in self.items]) + xyz_grid.axis_options.append(opt) + + def link_preview(self, filename): + quoted_filename = urllib.parse.quote(filename.replace('\\', '/')) + mtime = os.path.getmtime(filename) + preview = f"./sd_extra_networks/thumb?filename={quoted_filename}&mtime={mtime}" + return preview + + def search_terms_from_path(self, filename): + return filename.replace('\\', '/') + + def is_empty(self, folder): + return any(files_cache.list_files(folder, ext_filter=['.ckpt', '.safetensors', '.pt', '.json'])) + + def create_thumb(self): + debug(f'EN create-thumb: {self.name}') + created = 0 + for f in self.missing_thumbs: + if not os.path.exists(f): + continue + fn, _ext = os.path.splitext(f) + fn = fn.replace('.preview', '') + fn = f'{fn}.thumb.jpg' + if os.path.exists(fn): + continue + img = None + try: + img = Image.open(f) + except Exception: + img = None + shared.log.warning(f'Extra network removing invalid image: {f}') + try: + if img is None: + img = None + os.remove(f) + elif img.width > 1024 or img.height > 1024 or os.path.getsize(f) > 65536: + img = img.convert('RGB') + img.thumbnail((512, 512), Image.Resampling.HAMMING) + img.save(fn, quality=50) + img.close() + created += 1 + except Exception as e: + shared.log.warning(f'Extra network error creating thumbnail: {f} {e}') + if created > 0: + shared.log.info(f"Extra network thumbnails: {self.name} created={created}") + self.missing_thumbs.clear() + + def create_items(self, tabname): + if self.refresh_time is not None and self.refresh_time > refresh_time: # cached results + return + t0 = time.time() + try: + self.items = list(self.list_items()) + self.refresh_time = time.time() + except Exception as e: + self.items = [] + shared.log.error(f'Extra networks error listing items: class={self.__class__.__name__} tab={tabname} {e}') + for item in self.items: + if item is None: + continue + self.metadata[item["name"]] = item.get("metadata", {}) + t1 = time.time() + debug(f'EN create-items: page={self.name} items={len(self.items)} time={t1-t0:.2f}') + self.list_time += t1-t0 + + + def create_page(self, tabname, skip = False): + debug(f'EN create-page: {self.name}') + if self.page_time > refresh_time and len(self.html) > 0: # cached page + return self.html + self_name_id = self.name.replace(" ", "_") + if skip: + return f"
      Extra network page not ready
      Click refresh to try again
      " + subdirs = {} + allowed_folders = [os.path.abspath(x) for x in self.allowed_directories_for_previews()] + for parentdir, dirs in {d: files_cache.walk(d, cached=True, recurse=files_cache.not_hidden) for d in allowed_folders}.items(): + for tgt in dirs: + tgt = tgt.path + if os.path.join(paths.models_path, 'Reference') in tgt: + subdirs['Reference'] = 1 + if shared.backend == shared.Backend.DIFFUSERS and shared.opts.diffusers_dir in tgt: + subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 + if 'models--' in tgt: + continue + subdir = tgt[len(parentdir):].replace("\\", "/") + while subdir.startswith("/"): + subdir = subdir[1:] + if not subdir: + continue + # if not self.is_empty(tgt): + subdirs[subdir] = 1 + debug(f"Extra networks: page='{self.name}' subfolders={list(subdirs)}") + subdirs = OrderedDict(sorted(subdirs.items())) + if self.name == 'model': + subdirs['Reference'] = 1 + subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 + subdirs.move_to_end(os.path.basename(shared.opts.diffusers_dir)) + subdirs.move_to_end('Reference') + if self.name == 'style' and shared.opts.extra_networks_styles: + subdirs['built-in'] = 1 + subdirs_html = "
      " + subdirs_html += "".join([f"
      " for subdir in subdirs if subdir != '']) + self.html = '' + self.create_items(tabname) + self.create_xyz_grid() + htmls = [] + if len(self.items) > 0 and self.items[0].get('mtime', None) is not None: + self.items.sort(key=lambda x: x["mtime"], reverse=True) + for item in self.items: + htmls.append(self.create_html(item, tabname)) + self.html += ''.join(htmls) + self.page_time = time.time() + if len(subdirs_html) > 0 or len(self.html) > 0: + self.html = f"
      {subdirs_html}
      {self.html}
      " + else: + return '' + shared.log.debug(f"Extra networks: page='{self.name}' items={len(self.items)} subfolders={len(subdirs)} tab={tabname} folders={self.allowed_directories_for_previews()} list={self.list_time:.2f} desc={self.desc_time:.2f} info={self.info_time:.2f} workers={shared.max_workers}") + if len(self.missing_thumbs) > 0: + threading.Thread(target=self.create_thumb).start() + return self.html + + def list_items(self): + raise NotImplementedError + + def allowed_directories_for_previews(self): + return [] + + def create_html(self, item, tabname): + try: + args = { + "tabname": tabname, + "page": self.name, + "name": item["name"], + "title": os.path.basename(item["name"].replace('_', ' ')), + "filename": item["filename"], + "tags": '|'.join([item.get("tags")] if isinstance(item.get("tags", {}), str) else list(item.get("tags", {}).keys())), + "preview": html.escape(item.get("preview", self.link_preview('html/card-no-preview.png'))), + "width": shared.opts.extra_networks_card_size, + "height": shared.opts.extra_networks_card_size if shared.opts.extra_networks_card_square else 'auto', + "fit": shared.opts.extra_networks_card_fit, + "prompt": item.get("prompt", None), + "search": item.get("search_term", ""), + "description": item.get("description") or "", + "card_click": item.get("onclick", '"' + html.escape(f'return cardClicked({item.get("prompt", None)}, {"true" if self.allow_negative_prompt else "false"})') + '"'), + "mtime": item.get("mtime", 0), + "size": item.get("size", 0), + } + alias = item.get("alias", None) + if alias is not None: + args['title'] += f'\nAlias: {alias}' + return self.card.format(**args) + except Exception as e: + shared.log.error(f'Extra networks item error: page={tabname} item={item["name"]} {e}') + return "" + + def find_preview_file(self, path): + exts = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] + if path is None: + return 'html/card-no-preview.png' + if shared.opts.diffusers_dir in path: + path = os.path.relpath(path, shared.opts.diffusers_dir) + ref = os.path.join('models', 'Reference') + fn = os.path.join(ref, path.replace('models--', '').replace('\\', '/').split('/')[0]) + files = list(files_cache.list_files(ref, ext_filter=exts, recursive=False)) + else: + files = list(files_cache.list_files(os.path.dirname(path), ext_filter=exts, recursive=False)) + fn = os.path.splitext(path)[0] + for file in [f'{fn}{mid}{ext}' for ext in exts for mid in ['.thumb.', '.', '.preview.']]: + if file in files: + if 'Reference' not in file and '.thumb.' not in file: + self.missing_thumbs.append(file) + return file + return 'html/card-no-preview.png' + + def find_preview(self, path): + preview_file = self.find_preview_file(path) + return self.link_preview(preview_file) + + def find_description(self, path, info=None): + t0 = time.time() + class HTMLFilter(HTMLParser): + text = "" + def handle_data(self, data): + self.text += data + def handle_endtag(self, tag): + if tag == 'p': + self.text += '\n' + + fn = os.path.splitext(path)[0] + '.txt' + if os.path.exists(fn): + try: + with open(fn, "r", encoding="utf-8", errors="replace") as f: + txt = f.read() + txt = re.sub('[<>]', '', txt) + return txt + except OSError: + pass + if info is None: + info = self.find_info(path) + desc = info.get('description', '') or '' + f = HTMLFilter() + f.feed(desc) + t1 = time.time() + self.desc_time += t1-t0 + return f.text + + def find_info(self, path): + fn = os.path.splitext(path)[0] + '.json' + data = {} + if os.path.exists(fn): + t0 = time.time() + data = shared.readfile(fn, silent=True) + if type(data) is list: + data = data[0] + t1 = time.time() + self.info_time += t1-t0 + return data + + +def initialize(): + shared.extra_networks.clear() + + +def register_page(page: ExtraNetworksPage): + # registers extra networks page for the UI; recommend doing it in on_before_ui() callback for extensions + debug(f'EN register-page: {page}') + if page in shared.extra_networks: + debug(f'EN register-page: {page} already registered') + return + shared.extra_networks.append(page) + # allowed_dirs.clear() + # for pg in shared.extra_networks: + for folder in page.allowed_directories_for_previews(): + if folder not in allowed_dirs: + allowed_dirs.append(os.path.abspath(folder)) + + +def register_pages(): + from modules.ui_extra_networks_textual_inversion import ExtraNetworksPageTextualInversion + from modules.ui_extra_networks_hypernets import ExtraNetworksPageHypernetworks + from modules.ui_extra_networks_checkpoints import ExtraNetworksPageCheckpoints + from modules.ui_extra_networks_styles import ExtraNetworksPageStyles + from modules.ui_extra_networks_vae import ExtraNetworksPageVAEs + debug('EN register-pages') + register_page(ExtraNetworksPageCheckpoints()) + register_page(ExtraNetworksPageStyles()) + register_page(ExtraNetworksPageTextualInversion()) + register_page(ExtraNetworksPageHypernetworks()) + register_page(ExtraNetworksPageVAEs()) + + +def get_pages(title=None): + pages = [] + if 'All' in shared.opts.extra_networks: + pages = shared.extra_networks + else: + titles = [page.title for page in shared.extra_networks] + if title is None: + for page in shared.opts.extra_networks: + try: + idx = titles.index(page) + pages.append(shared.extra_networks[idx]) + except ValueError: + continue + else: + try: + idx = titles.index(title) + pages.append(shared.extra_networks[idx]) + except ValueError: + pass + return pages + + +class ExtraNetworksUi: + def __init__(self): + self.tabname: str = None + self.pages: list(str) = None + self.visible: gr.State = None + self.state: gr.Textbox = None + self.details: gr.Group = None + self.tabs: gr.Tabs = None + self.gallery: gr.Gallery = None + self.description: gr.Textbox = None + self.search: gr.Textbox = None + self.button_details: gr.Button = None + self.button_refresh: gr.Button = None + self.button_scan: gr.Button = None + self.button_view: gr.Button = None + self.button_quicksave: gr.Button = None + self.button_save: gr.Button = None + self.button_sort: gr.Button = None + self.button_apply: gr.Button = None + self.button_close: gr.Button = None + self.button_model: gr.Checkbox = None + self.details_components: list = [] + self.last_item: dict = None + self.last_page: ExtraNetworksPage = None + self.state: gr.State = None + + +def create_ui(container, button_parent, tabname, skip_indexing = False): + debug(f'EN create-ui: {tabname}') + ui = ExtraNetworksUi() + ui.tabname = tabname + ui.pages = [] + ui.state = gr.Textbox('{}', elem_id=f"{tabname}_extra_state", visible=False) + ui.visible = gr.State(value=False) # pylint: disable=abstract-class-instantiated + ui.details = gr.Group(elem_id=f"{tabname}_extra_details", visible=False) + ui.tabs = gr.Tabs(elem_id=f"{tabname}_extra_tabs") + ui.button_details = gr.Button('Details', elem_id=f"{tabname}_extra_details_btn", visible=False) + state = {} + if shared.cmd_opts.profile: + import cProfile + pr = cProfile.Profile() + pr.enable() + + def get_item(state, params = None): + if params is not None and type(params) == dict: + page = next(iter([x for x in get_pages() if x.title == 'Style']), None) + item = page.create_style(params) + else: + if state is None or not hasattr(state, 'page') or not hasattr(state, 'item'): + return None, None + page = next(iter([x for x in get_pages() if x.title == state.page]), None) + if page is None: + return None, None + item = next(iter([x for x in page.items if x["name"] == state.item]), None) + if item is None: + return page, None + item = SimpleNamespace(**item) + ui.last_item = item + ui.last_page = page + return page, item + + # main event that is triggered when js updates state text field with json values, used to communicate js -> python + def state_change(state_text): + try: + nonlocal state + state = SimpleNamespace(**json.loads(state_text)) + except Exception as e: + shared.log.error(f'Extra networks state error: {e}') + return + _page, _item = get_item(state) + # shared.log.debug(f'Extra network: op={state.op} page={page.title if page is not None else None} item={item.filename if item is not None else None}') + + def toggle_visibility(is_visible): + is_visible = not is_visible + return is_visible, gr.update(visible=is_visible), gr.update(variant=("secondary-down" if is_visible else "secondary")) + + with ui.details: + details_close = ToolButton(symbols.close, elem_id=f"{tabname}_extra_details_close", elem_classes=['extra-details-close']) + details_close.click(fn=lambda: gr.update(visible=False), inputs=[], outputs=[ui.details]) + with gr.Row(): + with gr.Column(scale=1): + text = gr.HTML('
      title
      ') + ui.details_components.append(text) + with gr.Column(scale=1): + img = gr.Image(value=None, show_label=False, interactive=False, container=False, show_download_button=False, show_info=False, elem_id=f"{tabname}_extra_details_img", elem_classes=['extra-details-img']) + ui.details_components.append(img) + with gr.Row(): + btn_save_img = gr.Button('Replace', elem_classes=['small-button']) + btn_delete_img = gr.Button('Delete', elem_classes=['small-button']) + with gr.Tabs(): + with gr.Tab('Description'): + desc = gr.Textbox('', show_label=False, lines=8, placeholder="Extra network description...") + ui.details_components.append(desc) + with gr.Row(): + btn_save_desc = gr.Button('Save', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_save_desc') + btn_delete_desc = gr.Button('Delete', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_delete_desc') + btn_close_desc = gr.Button('Close', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_close_desc') + btn_close_desc.click(fn=lambda: gr.update(visible=False), _js='refeshDetailsEN', inputs=[], outputs=[ui.details]) + with gr.Tab('Model metadata'): + info = gr.JSON({}, show_label=False) + ui.details_components.append(info) + with gr.Row(): + btn_save_info = gr.Button('Save', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_save_info') + btn_delete_info = gr.Button('Delete', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_delete_info') + btn_close_info = gr.Button('Close', elem_classes=['small-button'], elem_id=f'{tabname}_extra_details_close_info') + btn_close_info.click(fn=lambda: gr.update(visible=False), _js='refeshDetailsEN', inputs=[], outputs=[ui.details]) + with gr.Tab('Embedded metadata'): + meta = gr.JSON({}, show_label=False) + ui.details_components.append(meta) + + with ui.tabs: + def ui_tab_change(page): + scan_visible = page in ['Model', 'Lora', 'Hypernetwork', 'Embedding'] + save_visible = page in ['Style'] + model_visible = page in ['Model'] + return [gr.update(visible=scan_visible), gr.update(visible=save_visible), gr.update(visible=model_visible)] + + ui.button_refresh = ToolButton(symbols.refresh, elem_id=f"{tabname}_extra_refresh") + ui.button_scan = ToolButton(symbols.scan, elem_id=f"{tabname}_extra_scan", visible=True) + ui.button_quicksave = ToolButton(symbols.book, elem_id=f"{tabname}_extra_quicksave", visible=False) + ui.button_save = ToolButton(symbols.book, elem_id=f"{tabname}_extra_save", visible=False) + ui.button_sort = ToolButton(symbols.sort, elem_id=f"{tabname}_extra_sort", visible=True) + ui.button_view = ToolButton(symbols.view, elem_id=f"{tabname}_extra_view", visible=True) + ui.button_close = ToolButton(symbols.close, elem_id=f"{tabname}_extra_close", visible=True) + ui.button_model = ToolButton(symbols.refine, elem_id=f"{tabname}_extra_model", visible=True) + ui.search = gr.Textbox('', show_label=False, elem_id=f"{tabname}_extra_search", placeholder="Search...", elem_classes="textbox", lines=2, container=False) + ui.description = gr.Textbox('', show_label=False, elem_id=f"{tabname}_description", elem_classes="textbox", lines=2, interactive=False, container=False) + + if ui.tabname == 'txt2img': # refresh only once + global refresh_time # pylint: disable=global-statement + refresh_time = time.time() + if not skip_indexing: + threads = [] + for page in get_pages(): + if os.environ.get('SD_EN_DEBUG', None) is not None: + threads.append(threading.Thread(target=page.create_items, args=[ui.tabname])) + threads[-1].start() + else: + page.create_items(ui.tabname) + for thread in threads: + thread.join() + for page in get_pages(): + page.create_page(ui.tabname, skip_indexing) + with gr.Tab(page.title, id=page.title.lower().replace(" ", "_"), elem_classes="extra-networks-tab") as tab: + page_html = gr.HTML(page.html, elem_id=f'{tabname}{page.name}_extra_page', elem_classes="extra-networks-page") + ui.pages.append(page_html) + tab.select(ui_tab_change, _js="getENActivePage", inputs=[ui.button_details], outputs=[ui.button_scan, ui.button_save, ui.button_model]) + if shared.cmd_opts.profile: + errors.profile(pr, 'ExtraNetworks') + pr.disable() + # ui.tabs.change(fn=ui_tab_change, inputs=[], outputs=[ui.button_scan, ui.button_save]) + + def fn_save_img(image): + if ui.last_item is None or ui.last_item.local_preview is None: + return 'html/card-no-preview.png' + images = list(ui.gallery.temp_files) # gallery cannot be used as input component so looking at most recently registered temp files + if len(images) < 1: + shared.log.warning(f'Extra network no image: item={ui.last_item.name}') + return 'html/card-no-preview.png' + try: + images.sort(key=lambda f: os.path.getmtime(f), reverse=True) + image = Image.open(images[0]) + except Exception as e: + shared.log.error(f'Extra network error opening image: item={ui.last_item.name} {e}') + return 'html/card-no-preview.png' + fn_delete_img(image) + if image.width > 512 or image.height > 512: + image = image.convert('RGB') + image.thumbnail((512, 512), Image.Resampling.HAMMING) + try: + image.save(ui.last_item.local_preview, quality=50) + shared.log.debug(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}"') + except Exception as e: + shared.log.error(f'Extra network save image: item={ui.last_item.name} filename="{ui.last_item.local_preview}" {e}') + return image + + def fn_delete_img(_image): + preview_extensions = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] + fn = os.path.splitext(ui.last_item.filename)[0] + for file in [f'{fn}{mid}{ext}' for ext in preview_extensions for mid in ['.thumb.', '.preview.', '.']]: + if os.path.exists(file): + os.remove(file) + shared.log.debug(f'Extra network delete image: item={ui.last_item.name} filename="{file}"') + return 'html/card-no-preview.png' + + def fn_save_desc(desc): + if hasattr(ui.last_item, 'type') and ui.last_item.type == 'Style': + params = ui.last_page.parse_desc(desc) + if params is not None: + fn_save_info(params) + else: + fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' + with open(fn, 'w', encoding='utf-8') as f: + f.write(desc) + shared.log.debug(f'Extra network save desc: item={ui.last_item.name} filename="{fn}"') + return desc + + def fn_delete_desc(desc): + if ui.last_item is None: + return desc + if hasattr(ui.last_item, 'type') and ui.last_item.type == 'Style': + fn = os.path.splitext(ui.last_item.filename)[0] + '.json' + else: + fn = os.path.splitext(ui.last_item.filename)[0] + '.txt' + if os.path.exists(fn): + shared.log.debug(f'Extra network delete desc: item={ui.last_item.name} filename="{fn}"') + os.remove(fn) + return '' + return desc + + def fn_save_info(info): + fn = os.path.splitext(ui.last_item.filename)[0] + '.json' + shared.writefile(info, fn, silent=True) + shared.log.debug(f'Extra network save info: item={ui.last_item.name} filename="{fn}"') + return info + + def fn_delete_info(info): + if ui.last_item is None: + return info + fn = os.path.splitext(ui.last_item.filename)[0] + '.json' + if os.path.exists(fn): + shared.log.debug(f'Extra network delete info: item={ui.last_item.name} filename="{fn}"') + os.remove(fn) + return '' + return info + + btn_save_img.click(fn=fn_save_img, _js='closeDetailsEN', inputs=[img], outputs=[img]) + btn_delete_img.click(fn=fn_delete_img, _js='closeDetailsEN', inputs=[img], outputs=[img]) + btn_save_desc.click(fn=fn_save_desc, _js='closeDetailsEN', inputs=[desc], outputs=[desc]) + btn_delete_desc.click(fn=fn_delete_desc, _js='closeDetailsEN', inputs=[desc], outputs=[desc]) + btn_save_info.click(fn=fn_save_info, _js='closeDetailsEN', inputs=[info], outputs=[info]) + btn_delete_info.click(fn=fn_delete_info, _js='closeDetailsEN', inputs=[info], outputs=[info]) + + def show_details(text, img, desc, info, meta, params): + page, item = get_item(state, params) + if item is not None and hasattr(item, 'name'): + stat = os.stat(item.filename) if os.path.exists(item.filename) else None + desc = item.description + fullinfo = shared.readfile(os.path.splitext(item.filename)[0] + '.json', silent=True) + if 'modelVersions' in fullinfo: # sanitize massive objects + fullinfo['modelVersions'] = [] + info = fullinfo + meta = page.metadata.get(item.name, {}) or {} + if type(meta) is str: + try: + meta = json.loads(meta) + except Exception: + meta = {} + if ui.last_item.preview.startswith('data:'): + b64str = ui.last_item.preview.split(',',1)[1] + img = Image.open(io.BytesIO(base64.b64decode(b64str))) + elif hasattr(item, 'local_preview') and os.path.exists(item.local_preview): + img = item.local_preview + else: + img = page.find_preview_file(item.filename) + lora = '' + model = '' + style = '' + note = '' + if not os.path.exists(item.filename): + note = f'
      Target filename: {item.filename}' + if page.title == 'Model': + merge = len(list(meta.get('sd_merge_models', {}))) + if merge > 0: + model += f'Merge models{merge} recipes' + if meta.get('modelspec.architecture', None) is not None: + model += f''' + Architecture{meta.get('modelspec.architecture', 'N/A')} + Title{meta.get('modelspec.title', 'N/A')} + Resolution{meta.get('modelspec.resolution', 'N/A')} + ''' + if page.title == 'Lora': + try: + tags = getattr(item, 'tags', {}) + tags = [f'{name}:{tags[name]}' for i, name in enumerate(tags)] + tags = ' '.join(tags) + except Exception: + tags = '' + try: + triggers = ' '.join(info.get('tags', [])) + except Exception: + triggers = '' + lora = f''' + Model tags{tags} + User tags{triggers} + Base model{meta.get('ss_sd_model_name', 'N/A')} + Resolution{meta.get('ss_resolution', 'N/A')} + Training images{meta.get('ss_num_train_images', 'N/A')} + Comment{meta.get('ss_training_comment', 'N/A')} + ''' + if page.title == 'Style': + style = f''' + Name{item.name} + Description{item.description} + Preview Embedded{item.preview.startswith('data:')} + ''' + desc = f'Name: {os.path.basename(item.name)}\nDescription: {item.description}\nPrompt: {item.prompt}\nNegative: {item.negative}\nExtra: {item.extra}\n' + text = f''' +

      {item.name}

      + + + + + + + + + {lora} + {model} + {style} +
      Type{page.title}
      Alias{getattr(item, 'alias', 'N/A')}
      Filename{item.filename}
      Hash{getattr(item, 'hash', 'N/A')}
      Size{round(stat.st_size/1024/1024, 2) if stat is not None else 'N/A'} MB
      Last modified{datetime.fromtimestamp(stat.st_mtime) if stat is not None else 'N/A'}
      + {note} + ''' + return [text, img, desc, info, meta, gr.update(visible=item is not None)] + + def ui_refresh_click(title): + pages = [] + for page in get_pages(): + if title is None or title == '' or title == page.title or len(page.html) == 0: + page.page_time = 0 + page.refresh_time = 0 + page.refresh() + page.create_page(ui.tabname) + shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname}") + pages.append(page.html) + ui.search.update(value = ui.search.value) + return pages + + def ui_view_cards(title): + pages = [] + for page in get_pages(): + if title is None or title == '' or title == page.title or len(page.html) == 0: + shared.opts.extra_networks_view = page.view + page.view = 'gallery' if page.view == 'list' else 'list' + page.card = card_full if page.view == 'gallery' else card_list + page.html = '' + page.create_page(ui.tabname) + shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname} view={page.view}") + pages.append(page.html) + ui.search.update(value = ui.search.value) + return pages + + def ui_scan_click(title): + from modules import ui_models + if ui_models.search_metadata_civit is not None: + ui_models.search_metadata_civit(True, title) + return ui_refresh_click(title) + + def ui_save_click(): + from modules import generation_parameters_copypaste + filename = os.path.join(paths.data_path, "params.txt") + if os.path.exists(filename): + with open(filename, "r", encoding="utf8") as file: + prompt = file.read() + else: + prompt = '' + params = generation_parameters_copypaste.parse_generation_parameters(prompt) + res = show_details(text=None, img=None, desc=None, info=None, meta=None, params=params) + return res + + def ui_quicksave_click(name): + from modules import generation_parameters_copypaste + fn = os.path.join(paths.data_path, "params.txt") + if os.path.exists(fn): + with open(fn, "r", encoding="utf8") as file: + prompt = file.read() + else: + prompt = '' + params = generation_parameters_copypaste.parse_generation_parameters(prompt) + fn = os.path.join(shared.opts.styles_dir, os.path.splitext(name)[0] + '.json') + prompt = params.get('Prompt', '') + item = { + "name": name, + "description": '', + "prompt": prompt, + "negative": params.get('Negative prompt', ''), + "extra": '', + # "type": 'Style', + # "title": name, + # "filename": fn, + # "search_term": None, + # "preview": None, + # "local_preview": None, + } + shared.writefile(item, fn, silent=True) + if len(prompt) > 0: + shared.log.debug(f"Extra network quick save style: item={name} filename='{fn}'") + else: + shared.log.warning(f"Extra network quick save model: item={name} filename='{fn}' prompt is empty") + + def ui_sort_cards(msg): + shared.log.debug(f'Extra networks: {msg}') + return msg + + dummy = gr.State(value=False) # pylint: disable=abstract-class-instantiated + button_parent.click(fn=toggle_visibility, inputs=[ui.visible], outputs=[ui.visible, container, button_parent]) + ui.button_close.click(fn=toggle_visibility, inputs=[ui.visible], outputs=[ui.visible, container]) + ui.button_sort.click(fn=ui_sort_cards, _js='sortExtraNetworks', inputs=[ui.search], outputs=[ui.description]) + ui.button_view.click(fn=ui_view_cards, inputs=[ui.search], outputs=ui.pages) + ui.button_refresh.click(fn=ui_refresh_click, _js='getENActivePage', inputs=[ui.search], outputs=ui.pages) + ui.button_scan.click(fn=ui_scan_click, _js='getENActivePage', inputs=[ui.search], outputs=ui.pages) + ui.button_save.click(fn=ui_save_click, inputs=[], outputs=ui.details_components + [ui.details]) + ui.button_quicksave.click(fn=ui_quicksave_click, _js="() => prompt('Prompt name', '')", inputs=[ui.search], outputs=[]) + ui.button_details.click(show_details, _js="getCardDetails", inputs=ui.details_components + [dummy], outputs=ui.details_components + [ui.details]) + ui.state.change(state_change, inputs=[ui.state], outputs=[]) + return ui + + +def setup_ui(ui, gallery): + ui.gallery = gallery diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index 51dfba403..730cbd502 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -1,75 +1,75 @@ -import json -import os -import concurrent -from modules import shared, sd_hijack, sd_models, ui_extra_networks, files_cache -from modules.textual_inversion.textual_inversion import Embedding - - -class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): - def __init__(self): - super().__init__('Embedding') - self.allow_negative_prompt = True - self.embeddings = [] - - def refresh(self): - if sd_models.model_data.sd_model is None: - return - if shared.backend == shared.Backend.ORIGINAL: - sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True) - elif hasattr(sd_models.model_data.sd_model, 'embedding_db'): - sd_models.model_data.sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) - - def create_item(self, embedding: Embedding): - record = None - try: - path, _ext = os.path.splitext(embedding.filename) - tags = {} - if embedding.tag is not None: - tags[embedding.tag]=1 - name = os.path.splitext(embedding.basename)[0] - record = { - "type": 'Embedding', - "name": name, - "filename": embedding.filename, - "preview": self.find_preview(embedding.filename), - "search_term": self.search_terms_from_path(name), - "prompt": json.dumps(f" {os.path.splitext(embedding.name)[0]}"), - "local_preview": f"{path}.{shared.opts.samples_format}", - "tags": tags, - "mtime": os.path.getmtime(embedding.filename), - "size": os.path.getsize(embedding.filename), - } - record["info"] = self.find_info(embedding.filename) - record["description"] = self.find_description(embedding.filename, record["info"]) - except Exception as e: - shared.log.debug(f"Extra networks error: type=embedding file={embedding.filename} {e}") - return record - - def list_items(self): - if sd_models.model_data.sd_model is None: - self.embeddings = [ - Embedding(vec=0, name=os.path.basename(embedding_path), filename=embedding_path) - for embedding_path - in files_cache.list_files( - shared.opts.embeddings_dir, - ext_filter=['.pt', '.safetensors'], - recursive=files_cache.not_hidden - ) - ] - elif shared.backend == shared.Backend.ORIGINAL: - self.embeddings = list(sd_hijack.model_hijack.embedding_db.word_embeddings.values()) - elif hasattr(sd_models.model_data.sd_model, 'embedding_db'): - self.embeddings = list(sd_models.model_data.sd_model.embedding_db.word_embeddings.values()) - else: - self.embeddings = [] - self.embeddings = sorted(self.embeddings, key=lambda emb: emb.filename) - - with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: - future_items = {executor.submit(self.create_item, net): net for net in self.embeddings} - for future in concurrent.futures.as_completed(future_items): - item = future.result() - if item is not None: - yield item - - def allowed_directories_for_previews(self): - return list(sd_hijack.model_hijack.embedding_db.embedding_dirs) +import json +import os +import concurrent +from modules import shared, sd_hijack, sd_models, ui_extra_networks, files_cache +from modules.textual_inversion.textual_inversion import Embedding + + +class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): + def __init__(self): + super().__init__('Embedding') + self.allow_negative_prompt = True + self.embeddings = [] + + def refresh(self): + if sd_models.model_data.sd_model is None: + return + if shared.backend == shared.Backend.ORIGINAL: + sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings(force_reload=True) + elif hasattr(sd_models.model_data.sd_model, 'embedding_db'): + sd_models.model_data.sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) + + def create_item(self, embedding: Embedding): + record = None + try: + path, _ext = os.path.splitext(embedding.filename) + tags = {} + if embedding.tag is not None: + tags[embedding.tag]=1 + name = os.path.splitext(embedding.basename)[0] + record = { + "type": 'Embedding', + "name": name, + "filename": embedding.filename, + "preview": self.find_preview(embedding.filename), + "search_term": self.search_terms_from_path(name), + "prompt": json.dumps(f" {os.path.splitext(embedding.name)[0]}"), + "local_preview": f"{path}.{shared.opts.samples_format}", + "tags": tags, + "mtime": os.path.getmtime(embedding.filename), + "size": os.path.getsize(embedding.filename), + } + record["info"] = self.find_info(embedding.filename) + record["description"] = self.find_description(embedding.filename, record["info"]) + except Exception as e: + shared.log.debug(f"Extra networks error: type=embedding file={embedding.filename} {e}") + return record + + def list_items(self): + if sd_models.model_data.sd_model is None: + self.embeddings = [ + Embedding(vec=0, name=os.path.basename(embedding_path), filename=embedding_path) + for embedding_path + in files_cache.list_files( + shared.opts.embeddings_dir, + ext_filter=['.pt', '.safetensors'], + recursive=files_cache.not_hidden + ) + ] + elif shared.backend == shared.Backend.ORIGINAL: + self.embeddings = list(sd_hijack.model_hijack.embedding_db.word_embeddings.values()) + elif hasattr(sd_models.model_data.sd_model, 'embedding_db'): + self.embeddings = list(sd_models.model_data.sd_model.embedding_db.word_embeddings.values()) + else: + self.embeddings = [] + self.embeddings = sorted(self.embeddings, key=lambda emb: emb.filename) + + with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: + future_items = {executor.submit(self.create_item, net): net for net in self.embeddings} + for future in concurrent.futures.as_completed(future_items): + item = future.result() + if item is not None: + yield item + + def allowed_directories_for_previews(self): + return list(sd_hijack.model_hijack.embedding_db.embedding_dirs)