mirror of
https://github.com/vladmandic/automatic
synced 2026-08-26 15:16:01 +02:00
Cleanup, Ruff, Pylint
This commit is contained in:
+488
-488
@@ -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"<lora:{name}:{multiplier}>")
|
||||
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"<lora:{name}:{multiplier}>")
|
||||
if added:
|
||||
params["Prompt"] += "\n" + "".join(added)
|
||||
|
||||
|
||||
list_available_networks()
|
||||
|
||||
Submodule extensions-builtin/stable-diffusion-webui-rembg updated: 54b723361f...56a44f552e
+158
-158
@@ -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"<p>{self.commit_hash[:8]}</p><p>{datetime.fromtimestamp(self.commit_date).strftime('%a %b%d %Y %H:%M')}</p>"
|
||||
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"<p>{self.commit_hash[:8]}</p><p>{datetime.fromtimestamp(self.commit_date).strftime('%a %b%d %Y %H:%M')}</p>"
|
||||
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]}')
|
||||
|
||||
+39
-40
@@ -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({})
|
||||
cache_folders = DirectoryCache({})
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+196
-196
@@ -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 += "<error>"
|
||||
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 += "<error>"
|
||||
self.unload()
|
||||
shared.state.end()
|
||||
return res
|
||||
|
||||
+1342
-1342
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+826
-826
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user