From fe2e02073a934d7d97bbfd2582fea06b9e844567 Mon Sep 17 00:00:00 2001 From: Haoming Date: Mon, 4 Aug 2025 14:59:13 +0800 Subject: [PATCH] lint --- .../sd_forge_lora/extra_networks_lora.py | 9 ++- .../sd_forge_lora/lora_logger.py | 6 +- extensions-builtin/sd_forge_lora/network.py | 31 ++++---- extensions-builtin/sd_forge_lora/networks.py | 29 ++++--- extensions-builtin/sd_forge_lora/preload.py | 11 ++- .../sd_forge_lora/scripts/lora_script.py | 75 ++++++++++--------- .../sd_forge_lora/ui_edit_user_metadata.py | 47 ++++++------ .../sd_forge_lora/ui_extra_networks_lora.py | 10 +-- 8 files changed, 109 insertions(+), 109 deletions(-) diff --git a/extensions-builtin/sd_forge_lora/extra_networks_lora.py b/extensions-builtin/sd_forge_lora/extra_networks_lora.py index 33edf465..d1d0c0e2 100644 --- a/extensions-builtin/sd_forge_lora/extra_networks_lora.py +++ b/extensions-builtin/sd_forge_lora/extra_networks_lora.py @@ -1,15 +1,16 @@ -from modules import extra_networks, shared import networks +from modules import extra_networks, shared + class ExtraNetworkLora(extra_networks.ExtraNetwork): def __init__(self): - super().__init__('lora') + super().__init__("lora") self.errors = {} """mapping of network names to the number of errors the network had during operation""" - remove_symbols = str.maketrans('', '', ":,") + remove_symbols = str.maketrans("", "", ":,") def activate(self, p, params_list): additional = shared.opts.sd_lora @@ -53,7 +54,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): p.lora_hashes[item.mentioned_name.translate(self.remove_symbols)] = item.network_on_disk.shorthash if p.lora_hashes: - p.extra_generation_params["Lora hashes"] = ', '.join(f'{k}: {v}' for k, v in p.lora_hashes.items()) + p.extra_generation_params["Lora hashes"] = ", ".join(f"{k}: {v}" for k, v in p.lora_hashes.items()) def deactivate(self, p): if self.errors: diff --git a/extensions-builtin/sd_forge_lora/lora_logger.py b/extensions-builtin/sd_forge_lora/lora_logger.py index d51de297..57628d6e 100644 --- a/extensions-builtin/sd_forge_lora/lora_logger.py +++ b/extensions-builtin/sd_forge_lora/lora_logger.py @@ -1,6 +1,6 @@ -import sys import copy import logging +import sys class ColoredFormatter(logging.Formatter): @@ -27,7 +27,5 @@ logger.propagate = False if not logger.handlers: handler = logging.StreamHandler(sys.stdout) - handler.setFormatter( - ColoredFormatter("[%(name)s]-%(levelname)s: %(message)s") - ) + handler.setFormatter(ColoredFormatter("[%(name)s]-%(levelname)s: %(message)s")) logger.addHandler(handler) diff --git a/extensions-builtin/sd_forge_lora/network.py b/extensions-builtin/sd_forge_lora/network.py index 8f4002f9..45668e50 100644 --- a/extensions-builtin/sd_forge_lora/network.py +++ b/extensions-builtin/sd_forge_lora/network.py @@ -1,11 +1,11 @@ -import os import enum +import os -from modules import sd_models, cache, errors, hashes, shared - +from modules import cache, errors, hashes, sd_models, shared metadata_tags_order = {"ss_sd_model_name": 1, "ss_resolution": 2, "ss_clip_skip": 3, "ss_num_train_images": 10, "ss_tag_frequency": 20} + class SdVersion(enum.Enum): Unknown = 1 SD1 = 2 @@ -14,6 +14,7 @@ class SdVersion(enum.Enum): # SD3 = 5 Flux = 6 + class NetworkOnDisk: def __init__(self, name, filename): self.name = name @@ -28,7 +29,7 @@ class NetworkOnDisk: if self.is_safetensors: try: - self.metadata = cache.cached_data_for_file('safetensors-metadata', "lora/" + self.name, filename, read_metadata) + self.metadata = cache.cached_data_for_file("safetensors-metadata", "lora/" + self.name, filename, read_metadata) except Exception as e: errors.display(e, f"reading lora {filename}") @@ -39,28 +40,24 @@ class NetworkOnDisk: self.metadata = m - self.alias = self.metadata.get('ss_output_name', self.name) + self.alias = self.metadata.get("ss_output_name", self.name) self.hash = None self.shorthash = None - self.set_hash( - self.metadata.get('sshs_model_hash') or - hashes.sha256_from_cache(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or - '' - ) + self.set_hash(self.metadata.get("sshs_model_hash") or hashes.sha256_from_cache(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or "") self.sd_version = self.detect_version() def detect_version(self): - if str(self.metadata.get('modelspec.implementation', '')) == 'https://github.com/black-forest-labs/flux': + if str(self.metadata.get("modelspec.implementation", "")) == "https://github.com/black-forest-labs/flux": return SdVersion.Flux - elif str(self.metadata.get('modelspec.architecture', '')) == 'flux-1-dev/lora': + elif str(self.metadata.get("modelspec.architecture", "")) == "flux-1-dev/lora": return SdVersion.Flux - elif str(self.metadata.get('modelspec.architecture', '')) == 'stable-diffusion-xl-v1-base/lora': + elif str(self.metadata.get("modelspec.architecture", "")) == "stable-diffusion-xl-v1-base/lora": return SdVersion.SDXL - elif str(self.metadata.get('ss_base_model_version', '')).startswith('sdxl_'): + elif str(self.metadata.get("ss_base_model_version", "")).startswith("sdxl_"): return SdVersion.SDXL - elif str(self.metadata.get('modelspec.architecture', '')) == 'stable-diffusion-v1/lora': + elif str(self.metadata.get("modelspec.architecture", "")) == "stable-diffusion-v1/lora": return SdVersion.SD1 return SdVersion.Unknown @@ -71,14 +68,16 @@ class NetworkOnDisk: if self.shorthash: import networks + networks.available_network_hash_lookup[self.shorthash] = self def read_hash(self): if not self.hash: - self.set_hash(hashes.sha256(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or '') + self.set_hash(hashes.sha256(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or "") def get_alias(self): import networks + if shared.opts.lora_preferred_name == "Filename" or self.alias.lower() in networks.forbidden_network_aliases: return self.name else: diff --git a/extensions-builtin/sd_forge_lora/networks.py b/extensions-builtin/sd_forge_lora/networks.py index 5ff085b1..eb73c252 100644 --- a/extensions-builtin/sd_forge_lora/networks.py +++ b/extensions-builtin/sd_forge_lora/networks.py @@ -1,19 +1,20 @@ from __future__ import annotations +import functools import os import re -import torch -import network -import functools from typing import TYPE_CHECKING +import network +import torch + if TYPE_CHECKING: from backend.patcher.unet import UnetPatcher from backend.args import dynamic_args -from modules import shared, sd_models, errors, scripts +from backend.patcher.lora import load_lora, model_lora_keys_clip, model_lora_keys_unet from backend.utils import load_torch_file -from backend.patcher.lora import model_lora_keys_clip, model_lora_keys_unet, load_lora +from modules import errors, scripts, sd_models, shared def load_lora_for_models(model: "UnetPatcher", clip, lora, strength_model, strength_clip, filename="default", online_mode=False): @@ -40,11 +41,11 @@ def load_lora_for_models(model: "UnetPatcher", clip, lora, strength_model, stren lora_clip, lora_unmatch = load_lora(lora_unmatch, clip_keys) if len(lora_unmatch) > 12: - print(f'[LORA] LoRA version mismatch for {model_flag}: {filename}') + print(f"[LORA] LoRA version mismatch for {model_flag}: {filename}") return model, clip if len(lora_unmatch) > 0: - print(f'[LORA] Loading {filename} for {model_flag} with unmatched keys {list(lora_unmatch.keys())}') + print(f"[LORA] Loading {filename} for {model_flag} with unmatched keys {list(lora_unmatch.keys())}") new_model = model.clone() if model is not None else None new_clip = clip.clone() if clip is not None else None @@ -53,18 +54,18 @@ def load_lora_for_models(model: "UnetPatcher", clip, lora, strength_model, stren loaded_keys = new_model.add_patches(filename=filename, patches=lora_unet, strength_patch=strength_model, online_mode=online_mode) skipped_keys = [item for item in lora_unet if item not in loaded_keys] if len(skipped_keys) > 12: - print(f'[LORA] Mismatch {filename} for {model_flag}-UNet with {len(skipped_keys)} keys mismatched in {len(loaded_keys)} keys') + print(f"[LORA] Mismatch {filename} for {model_flag}-UNet with {len(skipped_keys)} keys mismatched in {len(loaded_keys)} keys") else: - print(f'[LORA] Loaded {filename} for {model_flag}-UNet with {len(loaded_keys)} keys at weight {strength_model} (skipped {len(skipped_keys)} keys) with on_the_fly = {online_mode}') + print(f"[LORA] Loaded {filename} for {model_flag}-UNet with {len(loaded_keys)} keys at weight {strength_model} (skipped {len(skipped_keys)} keys) with on_the_fly = {online_mode}") model = new_model if new_clip is not None and len(lora_clip) > 0: loaded_keys = new_clip.add_patches(filename=filename, patches=lora_clip, strength_patch=strength_clip, online_mode=online_mode) skipped_keys = [item for item in lora_clip if item not in loaded_keys] if len(skipped_keys) > 12: - print(f'[LORA] Mismatch {filename} for {model_flag}-CLIP with {len(skipped_keys)} keys mismatched in {len(loaded_keys)} keys') + print(f"[LORA] Mismatch {filename} for {model_flag}-CLIP with {len(skipped_keys)} keys mismatched in {len(loaded_keys)} keys") else: - print(f'[LORA] Loaded {filename} for {model_flag}-CLIP with {len(loaded_keys)} keys at weight {strength_clip} (skipped {len(skipped_keys)} keys) with on_the_fly = {online_mode}') + print(f"[LORA] Loaded {filename} for {model_flag}-CLIP with {len(loaded_keys)} keys at weight {strength_clip} (skipped {len(skipped_keys)} keys) with on_the_fly = {online_mode}") clip = new_clip return model, clip @@ -116,7 +117,7 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No network_on_disk.read_hash() loaded_networks.append(net) - online_mode = dynamic_args.get('online_lora', False) + online_mode = dynamic_args.get("online_lora", False) if current_sd.forge_objects.unet.model.storage_dtype in [torch.float32, torch.float16, torch.bfloat16]: online_mode = False @@ -136,9 +137,7 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No for filename, strength_model, strength_clip, online_mode in compiled_lora_targets: lora_sd = load_lora_state_dict(filename) - current_sd.forge_objects.unet, current_sd.forge_objects.clip = load_lora_for_models( - current_sd.forge_objects.unet, current_sd.forge_objects.clip, lora_sd, strength_model, strength_clip, - filename=filename, online_mode=online_mode) + current_sd.forge_objects.unet, current_sd.forge_objects.clip = load_lora_for_models(current_sd.forge_objects.unet, current_sd.forge_objects.clip, lora_sd, strength_model, strength_clip, filename=filename, online_mode=online_mode) current_sd.forge_objects_after_applying_lora = current_sd.forge_objects.shallow_copy() return diff --git a/extensions-builtin/sd_forge_lora/preload.py b/extensions-builtin/sd_forge_lora/preload.py index 763f9421..87c04758 100644 --- a/extensions-builtin/sd_forge_lora/preload.py +++ b/extensions-builtin/sd_forge_lora/preload.py @@ -1,8 +1,13 @@ -import os +import os.path + from modules import paths from modules.paths_internal import normalized_filepath def preload(parser): - parser.add_argument("--lora-dir", type=normalized_filepath, help="Path to directory with Lora networks.", default=os.path.join(paths.models_path, 'Lora')) - parser.add_argument("--lyco-dir-backcompat", type=normalized_filepath, help="Path to directory with LyCORIS networks (for backawards compatibility; can also use --lyco-dir).", default=os.path.join(paths.models_path, 'LyCORIS')) + parser.add_argument( + "--lora-dir", + type=normalized_filepath, + help="Path to directory with Lora networks.", + default=os.path.join(paths.models_path, "Lora"), + ) diff --git a/extensions-builtin/sd_forge_lora/scripts/lora_script.py b/extensions-builtin/sd_forge_lora/scripts/lora_script.py index 50864882..66e86104 100644 --- a/extensions-builtin/sd_forge_lora/scripts/lora_script.py +++ b/extensions-builtin/sd_forge_lora/scripts/lora_script.py @@ -1,14 +1,14 @@ import re +import extra_networks_lora import gradio as gr -from fastapi import FastAPI - +import lora # noqa import network import networks -import lora # noqa:F401 -import extra_networks_lora import ui_extra_networks_lora -from modules import script_callbacks, ui_extra_networks, extra_networks, shared +from fastapi import FastAPI + +from modules import extra_networks, script_callbacks, shared, ui_extra_networks def before_ui(): @@ -22,43 +22,44 @@ script_callbacks.on_before_ui(before_ui) script_callbacks.on_infotext_pasted(networks.infotext_pasted) -shared.options_templates.update(shared.options_section(('extra_networks', "Extra Networks"), { - "sd_lora": shared.OptionInfo("None", "Add network to prompt", gr.Dropdown, lambda: {"choices": ["None", *networks.available_networks]}, refresh=networks.list_available_networks), - "lora_preferred_name": shared.OptionInfo("Alias from file", "When adding to prompt, refer to Lora by", gr.Radio, {"choices": ["Alias from file", "Filename"]}), - "lora_add_hashes_to_infotext": shared.OptionInfo(True, "Add Lora hashes to infotext"), - "lora_bundled_ti_to_infotext": shared.OptionInfo(True, "Add Lora name as TI hashes for bundled Textual Inversion").info('"Add Textual Inversion hashes to infotext" needs to be enabled'), - "lora_filter_disabled": shared.OptionInfo(True, "Always show all networks on the Lora page").info("otherwise, those detected as for incompatible version of Stable Diffusion will be hidden"), - "lora_in_memory_limit": shared.OptionInfo(0, "Number of Lora networks to keep cached in memory", gr.Number, {"precision": 0}), - "lora_not_found_warning_console": shared.OptionInfo(False, "Lora not found warning in console"), - "lora_not_found_gradio_warning": shared.OptionInfo(False, "Lora not found warning popup in webui"), -})) +shared.options_templates.update( + shared.options_section( + ("extra_networks", "Extra Networks"), + { + "sd_lora": shared.OptionInfo("None", "Add network to prompt", gr.Dropdown, lambda: {"choices": ["None", *networks.available_networks]}, refresh=networks.list_available_networks), + "lora_preferred_name": shared.OptionInfo("Alias from file", "When adding to prompt, refer to Lora by", gr.Radio, {"choices": ["Alias from file", "Filename"]}), + "lora_add_hashes_to_infotext": shared.OptionInfo(True, "Add Lora hashes to infotext"), + "lora_bundled_ti_to_infotext": shared.OptionInfo(True, "Add Lora name as TI hashes for bundled Textual Inversion").info('"Add Textual Inversion hashes to infotext" needs to be enabled'), + "lora_filter_disabled": shared.OptionInfo(True, "Always show all networks on the Lora page").info("otherwise, those detected as for incompatible version of Stable Diffusion will be hidden"), + "lora_in_memory_limit": shared.OptionInfo(0, "Number of Lora networks to keep cached in memory", gr.Number, {"precision": 0}), + "lora_not_found_warning_console": shared.OptionInfo(False, "Lora not found warning in console"), + "lora_not_found_gradio_warning": shared.OptionInfo(False, "Lora not found warning popup in webui"), + }, + ) +) -shared.options_templates.update(shared.options_section(('compatibility', "Compatibility"), { - "lora_functional": shared.OptionInfo(False, "Lora/Networks: use old method that takes longer when you have multiple Loras active and produces same results as kohya-ss/sd-webui-additional-networks extension"), -})) +if shared.cmd_opts.api: + def create_lora_json(obj: network.NetworkOnDisk): + return { + "name": obj.name, + "alias": obj.alias, + "path": obj.filename, + "metadata": obj.metadata, + } -def create_lora_json(obj: network.NetworkOnDisk): - return { - "name": obj.name, - "alias": obj.alias, - "path": obj.filename, - "metadata": obj.metadata, - } + def api_networks(_: gr.Blocks, app: FastAPI): + @app.get("/sdapi/v1/loras") + async def get_loras(): + return [create_lora_json(obj) for obj in networks.available_networks.values()] + @app.post("/sdapi/v1/refresh-loras") + async def refresh_loras(): + return networks.list_available_networks() -def api_networks(_: gr.Blocks, app: FastAPI): - @app.get("/sdapi/v1/loras") - async def get_loras(): - return [create_lora_json(obj) for obj in networks.available_networks.values()] + script_callbacks.on_app_started(api_networks) - @app.post("/sdapi/v1/refresh-loras") - async def refresh_loras(): - return networks.list_available_networks() - - -script_callbacks.on_app_started(api_networks) re_lora = re.compile("{resolutions_text}" - table.append(('Resolutions:' if len(resolutions_list) > 1 else 'Resolution:', resolutions_text)) + table.append(("Resolutions:" if len(resolutions_list) > 1 else "Resolution:", resolutions_text)) image_count = 0 for _, params in metadata.get("ss_dataset_dirs", {}).items(): @@ -128,9 +128,9 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) *values[0:5], item.get("sd_version", "Unknown"), gr.HighlightedText.update(value=gradio_tags, visible=True if tags else False), - user_metadata.get('activation text', ''), - float(user_metadata.get('preferred weight', 0.0)), - user_metadata.get('negative text', ''), + user_metadata.get("activation text", ""), + float(user_metadata.get("preferred weight", 0.0)), + user_metadata.get("negative text", ""), gr.update(visible=True if tags else False), gr.update(value=self.generate_random_prompt_from_tags(tags), visible=True if tags else False), ] @@ -152,29 +152,29 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) v = random.random() * max_count if count > v: for x in "({[]})": - tag = tag.replace(x, '\\' + x) + tag = tag.replace(x, "\\" + x) res.append(tag) return ", ".join(sorted(res)) def create_extra_default_items_in_left_column(self): - self.select_sd_version = gr.Radio(['SD1', 'SDXL', 'Flux', 'Unknown'], value='Unknown', label='Base model', interactive=True) + self.select_sd_version = gr.Radio(["SD1", "SDXL", "Flux", "Unknown"], value="Unknown", label="Base model", interactive=True) def create_editor(self): self.create_default_editor_elems() self.taginfo = gr.HighlightedText(label="Training dataset tags") - self.edit_activation_text = gr.Text(label='Activation text', info="Will be added to prompt along with Lora") - self.slider_preferred_weight = gr.Slider(label='Preferred weight', info="Set to 0 to disable", minimum=0.0, maximum=2.0, step=0.01) - self.edit_negative_text = gr.Text(label='Negative prompt', info="Will be added to negative prompts") + self.edit_activation_text = gr.Text(label="Activation text", info="Will be added to prompt along with Lora") + self.slider_preferred_weight = gr.Slider(label="Preferred weight", info="Set to 0 to disable", minimum=0.0, maximum=2.0, step=0.01) + self.edit_negative_text = gr.Text(label="Negative prompt", info="Will be added to negative prompts") with gr.Row() as row_random_prompt: with gr.Column(scale=8): - random_prompt = gr.Textbox(label='Random prompt', lines=4, max_lines=4, interactive=False) + random_prompt = gr.Textbox(label="Random prompt", lines=4, max_lines=4, interactive=False) with gr.Column(scale=1, min_width=120): - generate_random_prompt = gr.Button('Generate', size="lg", scale=1) + generate_random_prompt = gr.Button("Generate", size="lg", scale=1) - self.edit_notes = gr.TextArea(label='Notes', lines=4) + self.edit_notes = gr.TextArea(label="Notes", lines=4) generate_random_prompt.click(fn=self.generate_random_prompt, inputs=[self.edit_name_input], outputs=[random_prompt], show_progress=False) @@ -207,9 +207,7 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) random_prompt, ] - self.button_edit\ - .click(fn=self.put_values_into_components, inputs=[self.edit_name_input], outputs=viewed_components)\ - .then(fn=lambda: gr.update(visible=True), inputs=[], outputs=[self.box]) + self.button_edit.click(fn=self.put_values_into_components, inputs=[self.edit_name_input], outputs=viewed_components).then(fn=lambda: gr.update(visible=True), inputs=[], outputs=[self.box]) edited_components = [ self.edit_description, @@ -220,5 +218,4 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) self.edit_notes, ] - self.setup_save_handler(self.button_save, self.save_lora_user_metadata, edited_components) diff --git a/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py b/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py index 9a649b1b..cccdabb8 100644 --- a/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py +++ b/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py @@ -2,15 +2,15 @@ import os import network import networks +from ui_edit_user_metadata import LoraUserMetadataEditor from modules import shared, ui_extra_networks from modules.ui_extra_networks import quote_js -from ui_edit_user_metadata import LoraUserMetadataEditor class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def __init__(self): - super().__init__('Lora') + super().__init__("Lora") self.allow_negative_prompt = True def refresh(self): @@ -37,7 +37,7 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): "search_terms": search_terms, "local_preview": f"{path}.{shared.opts.samples_format}", "metadata": lora_on_disk.metadata, - "sort_keys": {'default': index, **self.get_sort_keys(lora_on_disk.filename)}, + "sort_keys": {"default": index, **self.get_sort_keys(lora_on_disk.filename)}, } self.read_user_metadata(item) @@ -57,8 +57,8 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): item["sd_version"] = sd_version sd_version = network.SdVersion[sd_version] else: - sd_version = lora_on_disk.sd_version # use heuristics - #sd_version = network.SdVersion.Unknown # avoid heuristics + sd_version = lora_on_disk.sd_version # use heuristics + # sd_version = network.SdVersion.Unknown # avoid heuristics item["sd_version_str"] = str(sd_version)