This commit is contained in:
Haoming 2026-05-15 10:55:51 +08:00
parent cd837f4fed
commit 69dbcef1f8
4 changed files with 123 additions and 0 deletions

View File

@ -0,0 +1 @@
presets.json

View File

@ -0,0 +1,6 @@
import logging
from backend.logging import setup_logger
logger = logging.getLogger("Spectrum")
setup_logger(logger)

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

View File

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