From bce16a1db57978ce178bb78c68c3a97054876618 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 16 May 2023 12:24:04 -0400 Subject: [PATCH] update lora --- TODO.md | 6 +- .../Lora/extra_networks_lora.py | 1 + extensions-builtin/Lora/lora.py | 180 ++++++++++++------ extensions-builtin/Lora/preload.py | 6 + .../Lora/scripts/lora_script.py | 28 ++- .../Lora/ui_extra_networks_lora.py | 11 +- javascript/imageMaskFix.js | 2 +- modules/call_queue.py | 1 + modules/shared.py | 14 ++ modules/ui_extensions.py | 35 ++-- 10 files changed, 208 insertions(+), 76 deletions(-) create mode 100644 extensions-builtin/Lora/preload.py diff --git a/TODO.md b/TODO.md index 0c24fa5d7..566a10f0c 100644 --- a/TODO.md +++ b/TODO.md @@ -13,7 +13,6 @@ Stuff to be added... - Update `Wiki` - Add `Gradio` theme maker - Create new `GitHub` hooks/actions for CI/CD -- Redo Extensions tab: - Monitor file changes for misbehaving extensions - Kitchen theme: - Lightbox improvements @@ -40,7 +39,7 @@ Stuff to be investigated... Pick & merge PRs from main repo... -- Merge backlog: +- Merge backlog: ## Models @@ -66,4 +65,5 @@ Tech that can be integrated as part of the core workflow... ### Pending Code Updates -- Ability to save JSON log of all generated images with metadata +- add `--safe` mode which skips loading user extensions + please try to use it before opening new issue diff --git a/extensions-builtin/Lora/extra_networks_lora.py b/extensions-builtin/Lora/extra_networks_lora.py index 45f899fc4..ccb249ac7 100644 --- a/extensions-builtin/Lora/extra_networks_lora.py +++ b/extensions-builtin/Lora/extra_networks_lora.py @@ -1,6 +1,7 @@ from modules import extra_networks, shared import lora + class ExtraNetworkLora(extra_networks.ExtraNetwork): def __init__(self): super().__init__('lora') diff --git a/extensions-builtin/Lora/lora.py b/extensions-builtin/Lora/lora.py index 50dd59440..b5d0c98f9 100644 --- a/extensions-builtin/Lora/lora.py +++ b/extensions-builtin/Lora/lora.py @@ -1,10 +1,10 @@ import glob import os import re -from typing import Union import torch +from typing import Union -from modules import shared, devices, sd_models, errors +from modules import shared, devices, sd_models, errors, scripts metadata_tags_order = {"ss_sd_model_name": 1, "ss_resolution": 2, "ss_clip_skip": 3, "ss_num_train_images": 10, "ss_tag_frequency": 20} @@ -93,6 +93,7 @@ class LoraOnDisk: self.metadata = m self.ssmd_cover_images = self.metadata.pop('ssmd_cover_images', None) # those are cover images and they are too big to display in UI as text + self.alias = self.metadata.get('ss_output_name', self.name) class LoraModule: @@ -132,15 +133,17 @@ def load_lora(name, filename): sd = sd_models.read_state_dict(filename) + # this should not be needed but is here as an emergency fix for an unknown error people are experiencing in 1.2.0 + if not hasattr(shared.sd_model, 'lora_layer_mapping'): + assign_lora_names_to_compvis_modules(shared.sd_model) + keys_failed_to_match = {} is_sd2 = 'model_transformer_resblocks' in shared.sd_model.lora_layer_mapping - warnings = 0 for key_diffusers, weight in sd.items(): - lora_key_parts = key_diffusers.split(".", 1) - key_diffusers_without_lora_parts = lora_key_parts[0] - lora_key = lora_key_parts[1] if len(lora_key_parts) > 1 else "" + key_diffusers_without_lora_parts, lora_key = key_diffusers.split(".", 1) key = convert_diffusers_name_to_compvis(key_diffusers_without_lora_parts, is_sd2) + sd_module = shared.sd_model.lora_layer_mapping.get(key, None) if sd_module is None: @@ -167,13 +170,14 @@ def load_lora(name, filename): module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False) elif type(sd_module) == torch.nn.MultiheadAttention: module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False) - elif type(sd_module) == torch.nn.Conv2d: - module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (weight.shape[2], weight.shape[3]), bias=False) + elif type(sd_module) == torch.nn.Conv2d and weight.shape[2:] == (1, 1): + module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (1, 1), bias=False) + elif type(sd_module) == torch.nn.Conv2d and weight.shape[2:] == (3, 3): + module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (3, 3), bias=False) else: - if warnings == 0: - shared.log.warning(f'Lora layer {key_diffusers} matched a layer with unsupported type: {type(sd_module).__name__}') - warnings += 1 + print(f'Lora layer {key_diffusers} matched a layer with unsupported type: {type(sd_module).__name__}') continue + assert False, f'Lora layer {key_diffusers} matched a layer with unsupported type: {type(sd_module).__name__}' with torch.no_grad(): module.weight.copy_(weight) @@ -185,15 +189,10 @@ def load_lora(name, filename): elif lora_key == "lora_down.weight": lora_module.down = module else: - if warnings == 0: - shared.log.warning(f'Unknown Lora layer: {key_diffusers}') - shared.log.warning('Try using LyCORIS instead') - warnings += 1 + assert False, f'Bad Lora layer name: {key_diffusers} - must end in lora_up.weight, lora_down.weight or alpha' if len(keys_failed_to_match) > 0: - shared.log.warning(f"Lora failed to match keys when: {filename} {len(keys_failed_to_match)}") - shared.log.info(f"Try using LyCORIS instead") - warnings += 1 + print(f"Failed to match keys when loading Lora {filename}: {keys_failed_to_match}") return lora @@ -207,11 +206,11 @@ def load_loras(names, multipliers=None): loaded_loras.clear() - loras_on_disk = [available_loras.get(name, None) for name in names] + loras_on_disk = [available_lora_aliases.get(name, None) for name in names] if any([x is None for x in loras_on_disk]): list_available_loras() - loras_on_disk = [available_loras.get(name, None) for name in names] + loras_on_disk = [available_lora_aliases.get(name, None) for name in names] for i, name in enumerate(names): lora = already_loaded.get(name, None) @@ -226,7 +225,7 @@ def load_loras(names, multipliers=None): continue if lora is None: - shared.log.warning(f"Could not find Lora with name {name}") + print(f"Couldn't find Lora with name {name}") continue lora.multiplier = multipliers[i] if multipliers else 1.0 @@ -240,31 +239,29 @@ def lora_calc_updown(lora, module, target): if up.shape[2:] == (1, 1) and down.shape[2:] == (1, 1): updown = (up.squeeze(2).squeeze(2) @ down.squeeze(2).squeeze(2)).unsqueeze(2).unsqueeze(3) + elif up.shape[2:] == (3, 3) or down.shape[2:] == (3, 3): + updown = torch.nn.functional.conv2d(down.permute(1, 0, 2, 3), up).permute(1, 0, 2, 3) else: - permute, h, w = False, 1, 1 - if len(up.shape) == 4 and len(down.shape) == 4: - if up.shape[2:] == (1, 1): - up = up.squeeze(2).squeeze(2) - else: - n, c, h, w = up.shape - up = up.view(n, c, -1).permute(2, 0, 1) - permute = True - if down.shape[2:] == (1, 1): - down = down.squeeze(2).squeeze(2) - else: - n, c, h, w = down.shape - down = down.view(n, c, -1).permute(2, 0, 1) - permute = True updown = up @ down - if permute: - nh, nw = updown.shape[1:] - updown = updown.permute(1, 2, 0).view(nh, nw, h, w) updown = updown * lora.multiplier * (module.alpha / module.up.weight.shape[1] if module.alpha else 1.0) return updown +def lora_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.MultiheadAttention]): + weights_backup = getattr(self, "lora_weights_backup", None) + + if weights_backup is None: + return + + 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) + + def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.MultiheadAttention]): """ Applies the currently selected set of Loras to the weights of torch layer self. @@ -289,12 +286,7 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu self.lora_weights_backup = weights_backup if current_names != wanted_names: - 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) + lora_restore_weights_from_backup(self) for lora in loaded_loras: module = lora.modules.get(lora_layer_name, None) @@ -320,20 +312,53 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu if module is None: continue - shared.log.warning(f'failed to calculate lora weights for layer {lora_layer_name}') + print(f'failed to calculate lora weights for layer {lora_layer_name}') setattr(self, "lora_current_names", wanted_names) +def lora_forward(module, input, original_forward): + """ + 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_loras) == 0: + return original_forward(module, input) + + input = devices.cond_cast_unet(input) + + lora_restore_weights_from_backup(module) + lora_reset_cached_weight(module) + + res = original_forward(module, input) + + lora_layer_name = getattr(module, 'lora_layer_name', None) + for lora in loaded_loras: + module = lora.modules.get(lora_layer_name, None) + if module is None: + continue + + module.up.to(device=devices.device) + module.down.to(device=devices.device) + + res = res + module.up(module.down(input)) * lora.multiplier * (module.alpha / module.up.weight.shape[1] if module.alpha else 1.0) + + return res + + def lora_reset_cached_weight(self: Union[torch.nn.Conv2d, torch.nn.Linear]): setattr(self, "lora_current_names", ()) setattr(self, "lora_weights_backup", None) -def lora_Linear_forward(self, lora_input): +def lora_Linear_forward(self, input): + if shared.opts.lora_functional: + return lora_forward(self, input, torch.nn.Linear_forward_before_lora) + lora_apply_weights(self) - return torch.nn.Linear_forward_before_lora(self, lora_input) + return torch.nn.Linear_forward_before_lora(self, input) def lora_Linear_load_state_dict(self, *args, **kwargs): @@ -342,10 +367,13 @@ def lora_Linear_load_state_dict(self, *args, **kwargs): return torch.nn.Linear_load_state_dict_before_lora(self, *args, **kwargs) -def lora_Conv2d_forward(self, lora_input): +def lora_Conv2d_forward(self, input): + if shared.opts.lora_functional: + return lora_forward(self, input, torch.nn.Conv2d_forward_before_lora) + lora_apply_weights(self) - return torch.nn.Conv2d_forward_before_lora(self, lora_input) + return torch.nn.Conv2d_forward_before_lora(self, input) def lora_Conv2d_load_state_dict(self, *args, **kwargs): @@ -368,23 +396,65 @@ def lora_MultiheadAttention_load_state_dict(self, *args, **kwargs): def list_available_loras(): available_loras.clear() + available_lora_aliases.clear() + forbidden_lora_aliases.clear() + forbidden_lora_aliases.update({"none": 1}) + os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) - candidates = \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.pt'), recursive=True) + \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.safetensors'), recursive=True) + \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.ckpt'), recursive=True) - + candidates = list(shared.walk_files(shared.cmd_opts.lora_dir, allowed_extensions=[".pt", ".ckpt", ".safetensors"])) for filename in sorted(candidates, key=str.lower): if os.path.isdir(filename): continue name = os.path.splitext(os.path.basename(filename))[0] + entry = LoraOnDisk(name, filename) - available_loras[name] = LoraOnDisk(name, filename) + available_loras[name] = entry + if entry.alias in available_lora_aliases: + forbidden_lora_aliases[entry.alias.lower()] = 1 + + available_lora_aliases[name] = entry + available_lora_aliases[entry.alias] = entry + + +re_lora_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") + + +def infotext_pasted(infotext, params): + 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, v in params.items(): + 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_lora_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) available_loras = {} +available_lora_aliases = {} +forbidden_lora_aliases = {} loaded_loras = [] list_available_loras() diff --git a/extensions-builtin/Lora/preload.py b/extensions-builtin/Lora/preload.py new file mode 100644 index 000000000..863dc5c0b --- /dev/null +++ b/extensions-builtin/Lora/preload.py @@ -0,0 +1,6 @@ +import os +from modules import paths + + +def preload(parser): + parser.add_argument("--lora-dir", type=str, help="Path to directory with Lora networks.", default=os.path.join(paths.models_path, 'Lora')) diff --git a/extensions-builtin/Lora/scripts/lora_script.py b/extensions-builtin/Lora/scripts/lora_script.py index 3fc38ab9d..060bda059 100644 --- a/extensions-builtin/Lora/scripts/lora_script.py +++ b/extensions-builtin/Lora/scripts/lora_script.py @@ -1,12 +1,12 @@ import torch import gradio as gr +from fastapi import FastAPI import lora import extra_networks_lora import ui_extra_networks_lora from modules import script_callbacks, ui_extra_networks, extra_networks, shared - def unload(): torch.nn.Linear.forward = torch.nn.Linear_forward_before_lora torch.nn.Linear._load_from_state_dict = torch.nn.Linear_load_state_dict_before_lora @@ -49,8 +49,34 @@ torch.nn.MultiheadAttention._load_from_state_dict = lora.lora_MultiheadAttention script_callbacks.on_model_loaded(lora.assign_lora_names_to_compvis_modules) script_callbacks.on_script_unloaded(unload) script_callbacks.on_before_ui(before_ui) +script_callbacks.on_infotext_pasted(lora.infotext_pasted) shared.options_templates.update(shared.options_section(('extra_networks', "Extra Networks"), { "sd_lora": shared.OptionInfo("None", "Add Lora to prompt", gr.Dropdown, lambda: {"choices": ["None"] + [x for x in lora.available_loras]}, refresh=lora.list_available_loras), + "lora_preferred_name": shared.OptionInfo("Alias from file", "When adding to prompt, refer to lora by", gr.Radio, {"choices": ["Alias from file", "Filename"]}), })) + + +shared.options_templates.update(shared.options_section(('compatibility', "Compatibility"), { + "lora_functional": shared.OptionInfo(False, "Lora: use old method that takes longer when you have multiple Loras active and produces same results as kohya-ss/sd-webui-additional-networks extension"), +})) + + +def create_lora_json(obj: lora.LoraOnDisk): + return { + "name": obj.name, + "alias": obj.alias, + "path": obj.filename, + "metadata": obj.metadata, + } + + +def api_loras(_: gr.Blocks, app: FastAPI): + @app.get("/sdapi/v1/loras") + async def get_loras(): + return [create_lora_json(obj) for obj in lora.available_loras.values()] + + +script_callbacks.on_app_started(api_loras) + diff --git a/extensions-builtin/Lora/ui_extra_networks_lora.py b/extensions-builtin/Lora/ui_extra_networks_lora.py index c0a1ad2d3..2050e3faa 100644 --- a/extensions-builtin/Lora/ui_extra_networks_lora.py +++ b/extensions-builtin/Lora/ui_extra_networks_lora.py @@ -15,16 +15,23 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def list_items(self): for name, lora_on_disk in lora.available_loras.items(): path, ext = os.path.splitext(lora_on_disk.filename) + + if shared.opts.lora_preferred_name == "Filename" or lora_on_disk.alias.lower() in lora.forbidden_lora_aliases: + alias = name + else: + alias = lora_on_disk.alias + yield { "name": name, "filename": path, "preview": self.find_preview(path), "description": self.find_description(path), "search_term": self.search_terms_from_path(lora_on_disk.filename), - "prompt": json.dumps(f""), + "prompt": json.dumps(f""), "local_preview": f"{path}.{shared.opts.samples_format}", "metadata": json.dumps(lora_on_disk.metadata, indent=4) if lora_on_disk.metadata else None, } def allowed_directories_for_previews(self): - return [shared.opts.lora_dir] + return [shared.cmd_opts.lora_dir] + diff --git a/javascript/imageMaskFix.js b/javascript/imageMaskFix.js index 7c03fb394..ec64feaa6 100644 --- a/javascript/imageMaskFix.js +++ b/javascript/imageMaskFix.js @@ -2,7 +2,6 @@ * temporary fix for https://github.com/AUTOMATIC1111/stable-diffusion-webui/issues/668 * @see https://github.com/gradio-app/gradio/issues/1721 */ -window.addEventListener('resize', () => imageMaskResize()); function imageMaskResize() { const canvases = gradioApp().querySelectorAll('#img2maskimg .touch-none canvas'); if (!canvases.length) { @@ -42,4 +41,5 @@ function imageMaskResize() { }); } +window.addEventListener('resize', imageMaskResize); onUiUpdate(() => imageMaskResize()); diff --git a/modules/call_queue.py b/modules/call_queue.py index f5f4c935a..7fc2d4444 100644 --- a/modules/call_queue.py +++ b/modules/call_queue.py @@ -34,6 +34,7 @@ def wrap_gradio_gpu_call(func, extra_outputs=None): progress.record_results(id_task, res) except Exception as e: shared.log.error(f"Exception: {e}") + shared.log.error(f"Arguments: args={str(args)[:10240]} kwargs={str(kwargs)[:10240]}") errors.display(e, 'gradio call') res[-1] = f"
{html.escape(str(e))}
" finally: diff --git a/modules/shared.py b/modules/shared.py index 230f3f30c..afb3488a0 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -725,6 +725,20 @@ def listfiles(dirname): return [file for file in filenames if os.path.isfile(file)] +def walk_files(path, allowed_extensions=None): + if not os.path.exists(path): + return + if allowed_extensions is not None: + allowed_extensions = set(allowed_extensions) + for root, _dirs, files in os.walk(path): + for filename in files: + if allowed_extensions is not None: + _, ext = os.path.splitext(filename) + if ext not in allowed_extensions: + continue + yield os.path.join(root, filename) + + def html_path(filename): return os.path.join(paths.script_path, "html", filename) diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index fccc55693..23db404da 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -13,6 +13,19 @@ from modules.call_queue import wrap_gradio_gpu_call extensions_index = "https://vladmandic.github.io/sd-data/pages/extensions.json" hide_tags = ["localization"] extensions_list = [] +sort_ordering = { + "default": (True, lambda x: x.get('sort_string', '')), + "user extensions": (True, lambda x: x.get('sort_user', '')), + "update avilable": (True, lambda x: x.get('sort_update', '')), + "updated date": (True, lambda x: x.get('updated', '2000-01-01T00:00')), + "created date": (False, lambda x: x.get('created', '2000-01-01T00:00')), + "name": (False, lambda x: x.get('name', '').lower()), + "enabled": (False, lambda x: x.get('sort_enabled', '').lower()), + "size": (True, lambda x: x.get('size', 0)), + "stars": (True, lambda x: x.get('stars', 0)), + "commits": (True, lambda x: x.get('commits', 0)), + "issues": (True, lambda x: x.get('issues', 0)), +} def update_extension_list(): @@ -257,17 +270,6 @@ def refresh_extensions_list_from_data(search_text, sort_column): """ - sort_ordering = { - "default": (True, lambda x: x.get('sort_string', '')), - "updated": (True, lambda x: x.get('updated', '2000-01-01T00:00')), - "created": (False, lambda x: x.get('created', '2000-01-01T00:00')), - "name": (False, lambda x: x.get('name', '').lower()), - "enabled": (False, lambda x: x.get('sort_enabled', '').lower()), - "size": (True, lambda x: x.get('size', 0)), - "stars": (True, lambda x: x.get('stars', 0)), - "commits": (True, lambda x: x.get('commits', 0)), - "issues": (True, lambda x: x.get('issues', 0)), - } for ext in extensions_list: extension = [extension for extension in extensions.extensions if extension.git_name == ext['name'] or extension.name == ext['name']] if len(extension) > 0: @@ -279,8 +281,6 @@ def refresh_extensions_list_from_data(search_text, sort_column): ext['enabled'] = extension[0].enabled if len(extension) > 0 else '' ext['remote'] = extension[0].remote if len(extension) > 0 else None ext['path'] = extension[0].path if len(extension) > 0 else '' - ext['sort_string'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" - ext['sort_enabled'] = f"{'1' if ext['enabled'] else '0'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" sort_reverse, sort_function = sort_ordering[sort_column] def dt(x: str): @@ -308,6 +308,10 @@ def refresh_extensions_list_from_data(search_text, sort_column): remote = ext.get("remote", None) commit_date = ext.get("commit_date", 1577836800) or 1577836800 update_available = (remote is not None) & (installed) & (datetime.utcfromtimestamp(commit_date + 60 * 60) < datetime.fromisoformat(ext.get('updated', '2000-01-01T00:00:00.000Z')[:-1])) + ext['sort_string'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" + ext['sort_user'] = f"{'0' if ext['is_builtin'] else '1'}{'1' if ext['installed'] else '0'}{ext.get('name', '')}" + ext['sort_enabled'] = f"{'1' if ext['enabled'] else '0'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" + ext['sort_update'] = f"{'1' if update_available else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" tags = ext.get("tags", []) tags_string = ' '.join(tags) tags = tags + ["installed"] if installed else tags @@ -364,7 +368,10 @@ def create_ui(): search_text = gr.Text(label="Search") info = gr.HTML('Note: After any operation such as install/uninstall or enable/disable, please restart the server') with gr.Column(scale=1): - sort_column = gr.Dropdown(value="default", label="Sort by", choices=["default", "updated", "created", "name", "size", "stars", "commits", "issues"], multiselect=False) + print('HERE1', sort_ordering) + print('HERE2', list(sort_ordering.keys())) + print('HERE2', sort_ordering.items()) + sort_column = gr.Dropdown(value="default", label="Sort by", choices=list(sort_ordering.keys()), multiselect=False) with gr.Column(scale=1): refresh_extensions_button = gr.Button(value="Refresh extension list", variant="primary") check = gr.Button(value="Update installed extensions", variant="primary")