This commit is contained in:
Haoming 2025-08-04 14:59:13 +08:00
parent a6f71eb565
commit fe2e02073a
8 changed files with 109 additions and 109 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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