diff --git a/extensions-builtin/sd_forge_lora/network.py b/extensions-builtin/sd_forge_lora/network.py index 45668e50..f40d601c 100644 --- a/extensions-builtin/sd_forge_lora/network.py +++ b/extensions-builtin/sd_forge_lora/network.py @@ -1,69 +1,37 @@ -import enum import os +from typing import Final from modules import cache, errors, hashes, sd_models, shared +from modules_forge.presets import PresetArch -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 - # SD2 = 3 - SDXL = 4 - # SD3 = 5 - Flux = 6 +SD_VERSION: Final[list[str]] = ["Unknown"] + PresetArch.choices() class NetworkOnDisk: def __init__(self, name, filename): - self.name = name - self.filename = filename - self.metadata = {} - self.is_safetensors = os.path.splitext(filename)[1].lower() == ".safetensors" + self.name: str = name + self.filename: os.PathLike = filename + self.metadata: dict[str, str] = {} + self.is_safetensors: bool = filename.lower().endswith(".safetensors") def read_metadata(): metadata = sd_models.read_metadata_from_safetensors(filename) - return metadata if self.is_safetensors: try: 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}") + errors.display(e, f'reading metadata of "{filename}"') - if self.metadata: - m = {} - for k, v in sorted(self.metadata.items(), key=lambda x: metadata_tags_order.get(x[0], 999)): - m[k] = v + self.alias: str = self.metadata.get("ss_output_name", self.name) - self.metadata = m - - self.alias = self.metadata.get("ss_output_name", self.name) - - self.hash = None - self.shorthash = None + self.hash: bytes = None + self.shorthash: bytes = 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.sd_version = self.detect_version() - - def detect_version(self): - 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": - return SdVersion.Flux - 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_"): - return SdVersion.SDXL - elif str(self.metadata.get("modelspec.architecture", "")) == "stable-diffusion-v1/lora": - return SdVersion.SD1 - - return SdVersion.Unknown - - def set_hash(self, v): - self.hash = v + def set_hash(self, h): + self.hash = h self.shorthash = self.hash[0:12] if self.shorthash: @@ -75,7 +43,7 @@ class NetworkOnDisk: if not self.hash: self.set_hash(hashes.sha256(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or "") - def get_alias(self): + def get_alias(self) -> str: import networks if shared.opts.lora_preferred_name == "Filename" or self.alias.lower() in networks.forbidden_network_aliases: @@ -85,13 +53,11 @@ class NetworkOnDisk: class Network: - def __init__(self, name, network_on_disk: NetworkOnDisk): - self.name = name - self.network_on_disk = network_on_disk - self.te_multiplier = 1.0 - self.unet_multiplier = 1.0 - self.dyn_dim = None - self.modules = {} - self.bundle_embeddings = {} - self.mtime = None - self.mentioned_name = None + def __init__(self, name, network_on_disk): + self.name: str = name + self.network_on_disk: "NetworkOnDisk" = network_on_disk + self.te_multiplier: float = 1.0 + self.unet_multiplier: float = 1.0 + + self.mtime: float = None + self.mentioned_name: str = None diff --git a/extensions-builtin/sd_forge_lora/scripts/lora_script.py b/extensions-builtin/sd_forge_lora/scripts/lora_script.py index c282a28d..8cbc2bea 100644 --- a/extensions-builtin/sd_forge_lora/scripts/lora_script.py +++ b/extensions-builtin/sd_forge_lora/scripts/lora_script.py @@ -16,6 +16,7 @@ shared.options_templates.update( "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_preset_filter": shared.OptionInfo(False, "Filter Lora based on selected Preset"), }, ) ) diff --git a/extensions-builtin/sd_forge_lora/ui_edit_user_metadata.py b/extensions-builtin/sd_forge_lora/ui_edit_user_metadata.py index 905d4927..90bf9b36 100644 --- a/extensions-builtin/sd_forge_lora/ui_edit_user_metadata.py +++ b/extensions-builtin/sd_forge_lora/ui_edit_user_metadata.py @@ -8,9 +8,8 @@ import gradio as gr from modules import ui_extra_networks_user_metadata -def is_non_comma_tagset(tags): +def is_non_comma_tagset(tags: dict[str, int]) -> bool: average_tag_length = sum(len(x) for x in tags.keys()) / len(tags) - return average_tag_length >= 16 @@ -18,10 +17,10 @@ re_word = re.compile(r"[-_\w']+") re_comma = re.compile(r" *, *") -def build_tags(metadata): +def build_tags(metadata: dict) -> list[tuple[str, int]]: tags = {} - ss_tag_frequency = metadata.get("ss_tag_frequency", {}) + ss_tag_frequency: dict[str, dict[str, int]] = metadata.get("ss_tag_frequency", {}) if ss_tag_frequency is not None and hasattr(ss_tag_frequency, "items"): for _, tags_dict in ss_tag_frequency.items(): for tag, tag_count in tags_dict.items(): @@ -49,12 +48,12 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) def __init__(self, ui, tabname, page): super().__init__(ui, tabname, page) - self.select_sd_version = None + self.select_sd_version: gr.Dropdown = None - self.taginfo = None - self.edit_activation_text = None - self.slider_preferred_weight = None - self.edit_notes = None + self.taginfo: gr.HighlightedText = None + self.edit_activation_text: gr.Textbox = None + self.slider_preferred_weight: gr.Slider = None + self.edit_notes: gr.Textbox = None def save_lora_user_metadata(self, name, desc, sd_version, activation_text, preferred_weight, negative_text, notes): user_metadata = self.get_user_metadata(name) @@ -86,7 +85,7 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) ss_training_started_at = metadata.get("ss_training_started_at") if ss_training_started_at: - table.append(("Date trained:", datetime.datetime.utcfromtimestamp(float(ss_training_started_at)).strftime("%Y-%m-%d %H:%M"))) + table.append(("Date trained:", datetime.datetime.fromtimestamp(float(ss_training_started_at)).strftime("%Y-%m-%d %H:%M")), datetime.UTC) ss_bucket_info = metadata.get("ss_bucket_info") if ss_bucket_info and "buckets" in ss_bucket_info: @@ -158,7 +157,9 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor) 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) + import network + + self.select_sd_version = gr.Dropdown(choices=network.SD_VERSION, value="Unknown", label="Preset", interactive=True) def create_editor(self): self.create_default_editor_elems() 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 02df23d7..31a3d263 100644 --- a/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py +++ b/extensions-builtin/sd_forge_lora/ui_extra_networks_lora.py @@ -1,4 +1,4 @@ -import os +import os.path import network import networks @@ -21,14 +21,15 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): if lora_on_disk is None: return - path, ext = os.path.splitext(lora_on_disk.filename) + path = os.path.splitext(lora_on_disk.filename)[0] alias = lora_on_disk.get_alias() search_terms = [self.search_terms_from_path(lora_on_disk.filename)] if lora_on_disk.hash: search_terms.append(lora_on_disk.hash) - item = { + + item: dict[str, str | dict] = { "name": name, "filename": lora_on_disk.filename, "shorthash": lora_on_disk.shorthash, @@ -51,21 +52,18 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): negative_prompt = item["user_metadata"].get("negative text", "") item["negative_prompt"] = quote_js(negative_prompt) - # filter displayed loras by UI setting - sd_version = item["user_metadata"].get("sd version") - if sd_version in network.SdVersion.__members__: + sd_version: str = item["user_metadata"].get("sd version", None) + if sd_version in network.SD_VERSION: 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 = "Unknown" - item["sd_version_str"] = str(sd_version) + if enable_filter and shared.opts.lora_preset_filter and sd_version not in ("Unknown", shared.opts.forge_preset): + return None return item def list_items(self): - # instantiate a list to protect against concurrent modification names = list(networks.available_networks) for index, name in enumerate(names): item = self.create_item(name, index)