mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
lint
This commit is contained in:
parent
a6f71eb565
commit
fe2e02073a
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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"),
|
||||
)
|
||||
|
||||
@ -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("<lora:([^:]+):")
|
||||
|
||||
@ -68,7 +69,7 @@ def infotext_pasted(infotext, d):
|
||||
if not hashes:
|
||||
return
|
||||
|
||||
hashes = [x.strip().split(':', 1) for x in hashes.split(",")]
|
||||
hashes = [x.strip().split(":", 1) for x in hashes.split(",")]
|
||||
hashes = {x[0].strip().replace(",", ""): x[1].strip() for x in hashes}
|
||||
|
||||
def network_replacement(m):
|
||||
@ -81,7 +82,7 @@ def infotext_pasted(infotext, d):
|
||||
if network_on_disk is None:
|
||||
return m.group(0)
|
||||
|
||||
return f'<lora:{network_on_disk.get_alias()}:'
|
||||
return f"<lora:{network_on_disk.get_alias()}:"
|
||||
|
||||
d["Prompt"] = re.sub(re_lora, network_replacement, d["Prompt"])
|
||||
|
||||
|
||||
@ -1,9 +1,9 @@
|
||||
import datetime
|
||||
import html
|
||||
import random
|
||||
import re
|
||||
|
||||
import gradio as gr
|
||||
import re
|
||||
|
||||
from modules import ui_extra_networks_user_metadata
|
||||
|
||||
@ -22,7 +22,7 @@ def build_tags(metadata):
|
||||
tags = {}
|
||||
|
||||
ss_tag_frequency = metadata.get("ss_tag_frequency", {})
|
||||
if ss_tag_frequency is not None and hasattr(ss_tag_frequency, 'items'):
|
||||
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():
|
||||
tag = tag.strip()
|
||||
@ -73,10 +73,10 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor)
|
||||
metadata = item.get("metadata") or {}
|
||||
|
||||
keys = {
|
||||
'ss_output_name': "Output name:",
|
||||
'ss_sd_model_name': "Model:",
|
||||
'ss_clip_skip': "Clip skip:",
|
||||
'ss_network_module': "Kohya module:",
|
||||
"ss_output_name": "Output name:",
|
||||
"ss_sd_model_name": "Model:",
|
||||
"ss_clip_skip": "Clip skip:",
|
||||
"ss_network_module": "Kohya module:",
|
||||
}
|
||||
|
||||
for key, label in keys.items():
|
||||
@ -84,16 +84,16 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor)
|
||||
if value is not None and str(value) != "None":
|
||||
table.append((label, html.escape(value)))
|
||||
|
||||
ss_training_started_at = metadata.get('ss_training_started_at')
|
||||
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.utcfromtimestamp(float(ss_training_started_at)).strftime("%Y-%m-%d %H:%M")))
|
||||
|
||||
ss_bucket_info = metadata.get("ss_bucket_info")
|
||||
if ss_bucket_info and "buckets" in ss_bucket_info:
|
||||
resolutions = {}
|
||||
for _, bucket in ss_bucket_info["buckets"].items():
|
||||
resolution = bucket["resolution"]
|
||||
resolution = f'{resolution[1]}x{resolution[0]}'
|
||||
resolution = f"{resolution[1]}x{resolution[0]}"
|
||||
|
||||
resolutions[resolution] = resolutions.get(resolution, 0) + int(bucket["count"])
|
||||
|
||||
@ -103,7 +103,7 @@ class LoraUserMetadataEditor(ui_extra_networks_user_metadata.UserMetadataEditor)
|
||||
resolutions_text += ", ..."
|
||||
resolutions_text = f"<span title='{html.escape(', '.join(resolutions_list))}'>{resolutions_text}</span>"
|
||||
|
||||
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)
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user