This commit is contained in:
Haoming 2026-02-23 00:42:44 +08:00
parent 733298fb04
commit 2fa2e98155
4 changed files with 44 additions and 78 deletions

View File

@ -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

View File

@ -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"),
},
)
)

View File

@ -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()

View File

@ -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)