diff --git a/extensions-builtin/sd_forge_spectrum/.gitignore b/extensions-builtin/sd_forge_spectrum/.gitignore new file mode 100644 index 00000000..47eb2383 --- /dev/null +++ b/extensions-builtin/sd_forge_spectrum/.gitignore @@ -0,0 +1 @@ +presets.json diff --git a/extensions-builtin/sd_forge_spectrum/lib_spectrum/__init__.py b/extensions-builtin/sd_forge_spectrum/lib_spectrum/__init__.py new file mode 100644 index 00000000..dc547664 --- /dev/null +++ b/extensions-builtin/sd_forge_spectrum/lib_spectrum/__init__.py @@ -0,0 +1,6 @@ +import logging + +from backend.logging import setup_logger + +logger = logging.getLogger("Spectrum") +setup_logger(logger) diff --git a/extensions-builtin/sd_forge_spectrum/lib_spectrum/presets.py b/extensions-builtin/sd_forge_spectrum/lib_spectrum/presets.py new file mode 100644 index 00000000..654e2127 --- /dev/null +++ b/extensions-builtin/sd_forge_spectrum/lib_spectrum/presets.py @@ -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()) diff --git a/extensions-builtin/sd_forge_spectrum/scripts/spectrum.py b/extensions-builtin/sd_forge_spectrum/scripts/spectrum.py index 7f29b21a..dd3fe044 100644 --- a/extensions-builtin/sd_forge_spectrum/scripts/spectrum.py +++ b/extensions-builtin/sd_forge_spectrum/scripts/spectrum.py @@ -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]