mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
presets
This commit is contained in:
parent
cd837f4fed
commit
69dbcef1f8
1
extensions-builtin/sd_forge_spectrum/.gitignore
vendored
Normal file
1
extensions-builtin/sd_forge_spectrum/.gitignore
vendored
Normal file
@ -0,0 +1 @@
|
||||
presets.json
|
||||
@ -0,0 +1,6 @@
|
||||
import logging
|
||||
|
||||
from backend.logging import setup_logger
|
||||
|
||||
logger = logging.getLogger("Spectrum")
|
||||
setup_logger(logger)
|
||||
76
extensions-builtin/sd_forge_spectrum/lib_spectrum/presets.py
Normal file
76
extensions-builtin/sd_forge_spectrum/lib_spectrum/presets.py
Normal file
@ -0,0 +1,76 @@
|
||||
import os.path
|
||||
from json import dump, load
|
||||
from typing import Final
|
||||
|
||||
import gradio as gr
|
||||
|
||||
from lib_spectrum import logger
|
||||
|
||||
PRESET_FILE: Final[os.PathLike] = os.path.join(os.path.dirname(os.path.dirname(__file__)), "presets.json")
|
||||
PARAMS: Final[list[type]] = [float, int, float, int, float, int, float]
|
||||
|
||||
|
||||
class PresetManager:
|
||||
presets: dict[str, list[float]] = None
|
||||
|
||||
@classmethod
|
||||
def load_presets(cls):
|
||||
if cls.presets is not None:
|
||||
return
|
||||
|
||||
if not os.path.isfile(PRESET_FILE):
|
||||
with open(PRESET_FILE, "w+", encoding="utf-8") as json_file:
|
||||
dump({}, json_file)
|
||||
|
||||
logger.debug("Creating new empty Presets...")
|
||||
cls.presets = {}
|
||||
return
|
||||
|
||||
try:
|
||||
with open(PRESET_FILE, "r", encoding="utf-8") as json_file:
|
||||
cls.presets = load(json_file)
|
||||
except Exception:
|
||||
logger.error("Failed to load Presets...")
|
||||
cls.presets = {}
|
||||
else:
|
||||
logger.debug("Loaded Presets...")
|
||||
|
||||
@classmethod
|
||||
def list_preset(cls) -> list[str]:
|
||||
return list(cls.presets.keys())
|
||||
|
||||
@classmethod
|
||||
def get_preset(cls, preset_name: str) -> list[float]:
|
||||
if (preset := cls.presets.get(preset_name, None)) is None:
|
||||
logger.error(f'Preset "{preset_name}" was not found...')
|
||||
return [gr.skip()] * len(PARAMS)
|
||||
|
||||
return [gr.update(value=obj(val)) for obj, val in zip(PARAMS, preset)]
|
||||
|
||||
@classmethod
|
||||
def save_preset(cls, preset_name: str, *args: float) -> list[str]:
|
||||
if preset_name is None or not preset_name.strip():
|
||||
logger.error("Invalid Preset Name...")
|
||||
return gr.skip()
|
||||
|
||||
cls.presets.update({preset_name: [*args]})
|
||||
|
||||
with open(PRESET_FILE, "w", encoding="utf-8") as json_file:
|
||||
dump(cls.presets, json_file)
|
||||
|
||||
logger.info(f'Preset "{preset_name}" Saved!')
|
||||
return gr.update(choices=cls.list_preset())
|
||||
|
||||
@classmethod
|
||||
def delete_preset(cls, preset_name: str) -> list[str]:
|
||||
if preset_name not in cls.presets:
|
||||
logger.error(f'Preset "{preset_name}" was not found...')
|
||||
return gr.skip()
|
||||
|
||||
del cls.presets[preset_name]
|
||||
|
||||
with open(PRESET_FILE, "w", encoding="utf-8") as json_file:
|
||||
dump(cls.presets, json_file)
|
||||
|
||||
logger.info(f'Preset "{preset_name}" Deleted!')
|
||||
return gr.update(value=None, choices=cls.list_preset())
|
||||
@ -1,10 +1,13 @@
|
||||
import gradio as gr
|
||||
from lib_spectrum.forecaster import SpectrumNode
|
||||
from lib_spectrum.presets import PresetManager
|
||||
|
||||
from modules import scripts, shared
|
||||
from modules.infotext_utils import PasteField
|
||||
from modules.ui_components import InputAccordion
|
||||
|
||||
PresetManager.load_presets()
|
||||
|
||||
|
||||
class SpectrumForForge(scripts.Script):
|
||||
sorting_priority = 2026
|
||||
@ -77,6 +80,42 @@ class SpectrumForForge(scripts.Script):
|
||||
info="Run the full model for the last few steps",
|
||||
)
|
||||
|
||||
with gr.Accordion("Presets", open=False):
|
||||
_preset = gr.Dropdown(
|
||||
value=None,
|
||||
label="Preset Name",
|
||||
chois=PresetManager.list_preset(),
|
||||
allow_custom_value=True,
|
||||
)
|
||||
with gr.Row():
|
||||
_load = gr.Button("Apply Preset", variant="secondary")
|
||||
_save = gr.Button("Save Preset", variant="primary")
|
||||
_del = gr.Button("Delete Preset", variant="stop")
|
||||
|
||||
for comp in (_preset, _load, _save, _del):
|
||||
comp.do_not_save_to_config = True
|
||||
|
||||
args = (w, m, lam, window_size, flex_window, warmup_steps, stop_caching_step)
|
||||
|
||||
_load.click(
|
||||
fn=lambda name: PresetManager.get_preset(name),
|
||||
inputs=[_preset],
|
||||
outputs=[*args],
|
||||
queue=False,
|
||||
)
|
||||
_save.click(
|
||||
fn=lambda *args: PresetManager.save_preset(*args),
|
||||
inputs=[_preset, *args],
|
||||
outputs=[_preset],
|
||||
queue=False,
|
||||
)
|
||||
_del.click(
|
||||
fn=lambda name: PresetManager.delete_preset(name),
|
||||
inputs=[_preset],
|
||||
outputs=[_preset],
|
||||
queue=False,
|
||||
)
|
||||
|
||||
self.infotext_fields = [
|
||||
PasteField(w, "spec_w"),
|
||||
PasteField(m, "spec_m"),
|
||||
@ -86,6 +125,7 @@ class SpectrumForForge(scripts.Script):
|
||||
PasteField(warmup_steps, "spec_warmup_steps"),
|
||||
PasteField(stop_caching_step, "spec_stop_caching_step"),
|
||||
]
|
||||
self.paste_field_names = [field.label for field in self.infotext_fields]
|
||||
|
||||
return [enable, w, m, lam, window_size, flex_window, warmup_steps, stop_caching_step]
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user