mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
filter
This commit is contained in:
parent
733298fb04
commit
2fa2e98155
@ -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
|
||||
|
||||
@ -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"),
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user