From 56c382f71feb944de6d5b32ff5d7fa46d90ed7cd Mon Sep 17 00:00:00 2001 From: Mikei386 <44135113+Mikei386@users.noreply.github.com> Date: Tue, 8 Sep 2026 16:37:36 +0200 Subject: [PATCH] Persist ACE-Step XL-SFT quality defaults --- experiments/acestep15-xl-sft/README.md | 21 ++ experiments/acestep15-xl-sft/compose.yaml | 8 + .../generation_advanced_output_controls.py | 216 +++++++++++++++ .../overrides/model_config.py | 201 ++++++++++++++ .../overrides/user_preferences.js | 164 +++++++++++ .../overrides/user_preferences.py | 258 ++++++++++++++++++ 6 files changed, 868 insertions(+) create mode 100644 experiments/acestep15-xl-sft/overrides/generation_advanced_output_controls.py create mode 100644 experiments/acestep15-xl-sft/overrides/model_config.py create mode 100644 experiments/acestep15-xl-sft/overrides/user_preferences.js create mode 100644 experiments/acestep15-xl-sft/overrides/user_preferences.py diff --git a/experiments/acestep15-xl-sft/README.md b/experiments/acestep15-xl-sft/README.md index 17fc051..8ecc7df 100644 --- a/experiments/acestep15-xl-sft/README.md +++ b/experiments/acestep15-xl-sft/README.md @@ -9,6 +9,27 @@ Das offizielle Image ist auf den am 8. September 2026 geladenen Digest `sha256:95652cd780c78a1b1a7f6f0335530430f0ae53d96c7c12d59f9f39fa23d38567` fixiert. +## Persistente Qualitaetsvorgaben + +Die Weboberflaeche besitzt keine einzelne INI-Datei. Ihre Vorgaben kommen aus +Python-Modulen und teilweise aus dem Browser-`localStorage`. Deshalb bindet der +Compose-Dienst vier kleine, versionierte Overrides aus `./overrides` read-only +in den Container ein. Sie setzen fuer das XL-SFT-Modell: + +- 50 DiT-Schritte, Guidance 7, ODE/Euler und CFG-Intervall 0 bis 1 (Upstream-Defaults) +- Shift 1 statt des fehlerhaften UI-Werts 3 +- ADG aus, keine benutzerdefinierten Timesteps +- FLAC als verlustfreie Standardausgabe +- 320 kbit/s als MP3-Ausweichwert +- Batchgroesse 1 fuer einen einzelnen Qualitaetslauf +- Normalisierung an bei -1 dB, kein Fade, Latent Shift 0, Latent Rescale 1 + +Die Preference-Schema-Version wurde auf 2 angehoben. Alte, im Browser +gespeicherte MP3/128-kbit/s-Werte werden dadurch einmalig verworfen; danach +bleiben bewusst vorgenommene Aenderungen wieder im jeweiligen Browser erhalten. +Beim Wechsel des gepinnten Image-Digests muessen die Overrides gegen die neue +Upstream-Fassung geprueft werden. + ## Start Vor dem Start muessen das aktive llama.cpp-Profil und Qwen3-TTS beendet sein. diff --git a/experiments/acestep15-xl-sft/compose.yaml b/experiments/acestep15-xl-sft/compose.yaml index 1348755..ba7f2fa 100644 --- a/experiments/acestep15-xl-sft/compose.yaml +++ b/experiments/acestep15-xl-sft/compose.yaml @@ -11,6 +11,9 @@ services: ACESTEP_INIT_SERVICE: "true" ACESTEP_INIT_LLM: "true" ACESTEP_DEVICE: cuda + # The image entrypoint forwards only ACESTEP_EXTRA_ARGS to the UI CLI. + # One result per run avoids the batch=2 VRAM/time penalty. + ACESTEP_EXTRA_ARGS: "--batch_size 1" TOKENIZERS_PARALLELISM: "false" NVIDIA_VISIBLE_DEVICES: ${ACESTEP_GPU_UUID:?set ACESTEP_GPU_UUID to the RTX 5080 UUID} deploy: @@ -27,6 +30,11 @@ services: - ${ACESTEP_HF_CACHE_DIR:-/data/models/acestep/hf-cache}:/root/.cache/huggingface - ${ACESTEP_OUTPUT_DIR:-/data/music/acestep}:/app/gradio_outputs - ${ACESTEP_OUTPUT_DIR:-/data/music/acestep}:/app/output + # Version-pinned UI defaults for XL-SFT quality and lossless output. + - ./overrides/model_config.py:/app/acestep/ui/gradio/events/generation/model_config.py:ro + - ./overrides/generation_advanced_output_controls.py:/app/acestep/ui/gradio/interfaces/generation_advanced_output_controls.py:ro + - ./overrides/user_preferences.py:/app/acestep/ui/gradio/interfaces/user_preferences.py:ro + - ./overrides/user_preferences.js:/app/acestep/ui/gradio/interfaces/user_preferences.js:ro shm_size: "2gb" restart: "no" healthcheck: diff --git a/experiments/acestep15-xl-sft/overrides/generation_advanced_output_controls.py b/experiments/acestep15-xl-sft/overrides/generation_advanced_output_controls.py new file mode 100644 index 0000000..1229a9a --- /dev/null +++ b/experiments/acestep15-xl-sft/overrides/generation_advanced_output_controls.py @@ -0,0 +1,216 @@ +"""Output and automation controls for generation advanced settings.""" + +from typing import Any + +import gradio as gr + +from acestep.ui.gradio.i18n import t + + +_MP3_BITRATE_CHOICES = [("128 kbps", "128k"), ("192 kbps", "192k"), ("256 kbps", "256k"), ("320 kbps", "320k")] +_MP3_SAMPLE_RATE_CHOICES = [("48 kHz", 48000), ("44.1 kHz", 44100)] + + +def _update_mp3_control_visibility(audio_format: str, service_mode: bool = False): + """Return visibility and interactivity updates for MP3-only controls.""" + visible = audio_format == "mp3" + interactive = visible and not service_mode + return ( + gr.update(visible=visible), + gr.update(choices=_MP3_BITRATE_CHOICES, visible=visible, interactive=interactive), + gr.update(choices=_MP3_SAMPLE_RATE_CHOICES, visible=visible, interactive=interactive), + ) + + +def build_output_controls( + service_pre_initialized: bool, + service_mode: bool, + init_params: dict[str, Any] | None, +) -> dict[str, Any]: + """Create audio-output and post-processing controls for advanced settings. + + Args: + service_pre_initialized: Whether existing init params should prefill values. + service_mode: Whether the UI is running in service mode (disables some controls). + init_params: Optional startup state containing persisted output values. + + Returns: + A component map containing format, scoring, normalization, and latent controls. + """ + + params = init_params or {} + # Keep the master lossless. MP3 is only an optional sharing export. + initial_audio_format = params.get("audio_format", "flac") + initial_mp3_visible = initial_audio_format == "mp3" + with gr.Accordion(t("generation.advanced_output_section"), open=False, elem_classes=["has-info-container"]): + with gr.Row(): + with gr.Column(scale=1): + audio_format = gr.Dropdown( + choices=[ + ("FLAC", "flac"), + ("MP3", "mp3"), + ("Opus", "opus"), + ("AAC", "aac"), + ("WAV (16-bit)", "wav"), + ("WAV (32-bit Float)", "wav32"), + ], + value=initial_audio_format, + label=t("generation.audio_format_label"), + info=t("generation.audio_format_info"), + elem_id="acestep-audio-format", + elem_classes=["has-info-container"], + interactive=not service_mode, + ) + with gr.Row(visible=initial_mp3_visible) as mp3_controls_row: + mp3_bitrate = gr.Dropdown( + choices=[ + ("128 kbps", "128k"), + ("192 kbps", "192k"), + ("256 kbps", "256k"), + ("320 kbps", "320k"), + ], + value=params.get("mp3_bitrate", "320k"), + label=t("generation.mp3_bitrate_label"), + info=t("generation.mp3_bitrate_info"), + elem_id="acestep-mp3-bitrate", + elem_classes=["has-info-container"], + visible=initial_mp3_visible, + interactive=initial_mp3_visible and not service_mode, + scale=1, + ) + mp3_sample_rate = gr.Dropdown( + choices=[ + ("48 kHz", 48000), + ("44.1 kHz", 44100), + ], + value=params.get("mp3_sample_rate", 48000), + label=t("generation.mp3_sample_rate_label"), + info=t("generation.mp3_sample_rate_info"), + elem_id="acestep-mp3-sample-rate", + elem_classes=["has-info-container"], + visible=initial_mp3_visible, + interactive=initial_mp3_visible and not service_mode, + scale=1, + ) + with gr.Column(scale=1): + score_scale = gr.Slider( + minimum=0.01, + maximum=1.0, + value=0.5, + step=0.01, + label=t("generation.score_sensitivity_label"), + info=t("generation.score_sensitivity_info"), + elem_id="acestep-score-scale", + elem_classes=["has-info-container"], + scale=1, + visible=not service_mode, + ) + audio_format.change( + fn=lambda value: _update_mp3_control_visibility(value, service_mode), + inputs=[audio_format], + outputs=[mp3_controls_row, mp3_bitrate, mp3_sample_rate], + ) + with gr.Row(): + enable_normalization = gr.Checkbox( + label=t("generation.enable_normalization"), + value=params.get("enable_normalization", True) if service_pre_initialized else True, + info=t("generation.enable_normalization_info"), + elem_id="acestep-enable-normalization", + elem_classes=["has-info-container"], + ) + normalization_db = gr.Slider( + label=t("generation.normalization_db"), + minimum=-10.0, + maximum=0.0, + step=0.1, + value=params.get("normalization_db", -1.0) if service_pre_initialized else -1.0, + info=t("generation.normalization_db_info"), + elem_id="acestep-normalization-db", + elem_classes=["has-info-container"], + ) + with gr.Row(): + fade_in_duration = gr.Slider( + label=t("generation.fade_in_duration"), + minimum=0.0, + maximum=10.0, + step=0.1, + value=params.get("fade_in_duration", 0.0) if service_pre_initialized else 0.0, + info=t("generation.fade_in_duration_info"), + elem_id="acestep-fade-in-duration", + elem_classes=["has-info-container"], + ) + fade_out_duration = gr.Slider( + label=t("generation.fade_out_duration"), + minimum=0.0, + maximum=10.0, + step=0.1, + value=params.get("fade_out_duration", 0.0) if service_pre_initialized else 0.0, + info=t("generation.fade_out_duration_info"), + elem_id="acestep-fade-out-duration", + elem_classes=["has-info-container"], + ) + with gr.Row(): + latent_shift = gr.Slider( + label=t("generation.latent_shift"), + minimum=-0.2, + maximum=0.2, + step=0.01, + value=params.get("latent_shift", 0.0) if service_pre_initialized else 0.0, + info=t("generation.latent_shift_info"), + elem_id="acestep-latent-shift", + elem_classes=["has-info-container"], + ) + latent_rescale = gr.Slider( + label=t("generation.latent_rescale"), + minimum=0.5, + maximum=1.5, + step=0.01, + value=params.get("latent_rescale", 1.0) if service_pre_initialized else 1.0, + info=t("generation.latent_rescale_info"), + elem_id="acestep-latent-rescale", + elem_classes=["has-info-container"], + ) + return { + "audio_format": audio_format, + "mp3_controls_row": mp3_controls_row, + "mp3_bitrate": mp3_bitrate, + "mp3_sample_rate": mp3_sample_rate, + "score_scale": score_scale, + "enable_normalization": enable_normalization, + "normalization_db": normalization_db, + "fade_in_duration": fade_in_duration, + "fade_out_duration": fade_out_duration, + "latent_shift": latent_shift, + "latent_rescale": latent_rescale, + } + + +def build_automation_controls(service_mode: bool) -> dict[str, Any]: + """Create automation controls for LM batch chunking. + + Args: + service_mode: Whether the UI is running in service mode (disables some controls). + + Returns: + A component map containing ``lm_batch_chunk_size``. + """ + + with gr.Accordion( + t("generation.advanced_automation_section"), + open=False, + elem_classes=["has-info-container"], + ): + with gr.Row(): + lm_batch_chunk_size = gr.Number( + label=t("generation.lm_batch_chunk_label"), + value=8, + minimum=1, + maximum=32, + step=1, + info=t("generation.lm_batch_chunk_info"), + scale=1, + interactive=not service_mode, + elem_id="acestep-lm-batch-chunk-size", + elem_classes=["has-info-container"], + ) + return {"lm_batch_chunk_size": lm_batch_chunk_size} diff --git a/experiments/acestep15-xl-sft/overrides/model_config.py b/experiments/acestep15-xl-sft/overrides/model_config.py new file mode 100644 index 0000000..ed75ed8 --- /dev/null +++ b/experiments/acestep15-xl-sft/overrides/model_config.py @@ -0,0 +1,201 @@ +"""Model configuration and UI control settings for generation handlers. + +Contains functions for determining model type (turbo/base/pure-base), +producing UI control configurations, and computing gr.update() tuples +for model-type-dependent controls. +""" + +import re + +import gradio as gr + +from acestep.constants import ( + TASK_TYPES_TURBO, + TASK_TYPES_BASE, + GENERATION_MODES_TURBO, + GENERATION_MODES_BASE, +) + + +def _has_token(token: str, path: str) -> bool: + """Check if *token* appears as a delimited word in *path*. + + Matches when *token* is bounded by start/end of string or a common + path delimiter (``/``, ``\\``, ``.``, ``_``, ``-``). + """ + return re.search(rf"(^|[\\\\/._-]){token}($|[\\\\/._-])", path) is not None + + +def is_pure_base_model(config_path_lower: str) -> bool: + """Check whether a model path refers to a pure base model. + + Args: + config_path_lower: Lowercased model config path string. + + Returns: + ``True`` when the path contains ``"base"`` and excludes ``"sft"`` and ``"turbo"``. + """ + return ( + _has_token("base", config_path_lower) + and not _has_token("sft", config_path_lower) + and not _has_token("turbo", config_path_lower) + ) + + +def update_model_type_settings(config_path: str | None, current_mode: str | None = None) -> tuple: + """Update UI settings based on model type (fallback when handler not initialized yet). + + Args: + config_path: Model config path string. + current_mode: Current generation mode value to preserve across choices update. + + Returns: + Ten-element tuple of ``gr.update()`` dicts for inference_steps, + guidance_scale, use_adg, shift, cfg_interval_start, cfg_interval_end, + task_type, generation_mode, init_llm_checkbox, and dcw_enabled. + """ + if config_path is None: + config_path = "" + config_path_lower = config_path.lower() + + # Precedence: turbo > SFT > pure base > fallback. + # Detection functions enforce mutual exclusivity. + is_turbo = _has_token("turbo", config_path_lower) + is_pure_base = is_pure_base_model(config_path_lower) + is_sft = is_sft_model(config_path_lower) + + return get_model_type_ui_settings(is_turbo, current_mode=current_mode, is_pure_base=is_pure_base, is_sft=is_sft) + + +def is_sft_model(config_path_lower: str) -> bool: + """Check whether a model path refers to an SFT (supervised fine-tuned) model. + + Args: + config_path_lower: Lowercased model config path string. + + Returns: + ``True`` when the path contains ``"sft"`` and excludes ``"turbo"``. + """ + return _has_token("sft", config_path_lower) and not _has_token("turbo", config_path_lower) + + +def is_xl_model(config_path_lower: str) -> bool: + """Check whether a model path refers to an XL (4B DiT) variant. + + Args: + config_path_lower: Lowercased model config path string. + + Returns: + ``True`` when the path contains ``"xl"`` as a delimited token. + """ + return _has_token("xl", config_path_lower) + + +def get_ui_control_config(is_turbo: bool, is_pure_base: bool = False, is_sft: bool = False) -> dict: + """Return UI control configuration (values, limits, visibility) for model type. + + Args: + is_turbo: Whether the model is a turbo variant. + is_pure_base: Whether the model is a pure base model. + is_sft: Whether the model is an SFT (supervised fine-tuned) variant. + SFT models are optimized for 50 inference steps, matching the + training defaults in model_discovery._BASE_DEFAULTS. + + Used by both interactive init and service-mode startup so controls stay consistent. + """ + # Precedence: turbo > SFT > pure base > fallback. + if is_pure_base: + task_choices = TASK_TYPES_BASE + mode_choices = GENERATION_MODES_BASE + else: + task_choices = TASK_TYPES_TURBO + mode_choices = GENERATION_MODES_TURBO + + if is_turbo: + return { + "inference_steps_value": 8, + "inference_steps_maximum": 20, + "inference_steps_minimum": 1, + "guidance_scale_visible": False, + "use_adg_visible": False, + "shift_value": 3.0, + "shift_visible": True, + "dcw_enabled_value": True, + "cfg_interval_start_visible": False, + "cfg_interval_end_visible": False, + "task_type_choices": task_choices, + "generation_mode_choices": mode_choices, + } + else: + # SFT models use 50 steps; pure base / unknown models use 32. + steps = 50 if is_sft else 32 + return { + "inference_steps_value": steps, + "inference_steps_maximum": 200, + "inference_steps_minimum": 1, + "guidance_scale_visible": True, + "use_adg_visible": True, + # ACE-Step XL-SFT was trained/recommended with shift=1.0. + # Keep 3.0 only for non-SFT base/unknown variants. + "shift_value": 1.0 if is_sft else 3.0, + "shift_visible": True, + "dcw_enabled_value": False, + "cfg_interval_start_visible": True, + "cfg_interval_end_visible": True, + "task_type_choices": task_choices, + "generation_mode_choices": mode_choices, + } + + +def get_model_type_ui_settings(is_turbo: bool, current_mode: str | None = None, is_pure_base: bool = False, is_sft: bool = False): + """Get gr.update() tuple for model-type controls. + + Args: + is_turbo: Whether the model is a turbo variant. + current_mode: Current generation mode value to preserve. + is_pure_base: Whether the model is a pure base model. + is_sft: Whether the model is an SFT variant. + + Returns: + Tuple of updates for inference_steps, guidance_scale, use_adg, + shift, cfg_interval_start, cfg_interval_end, task_type, + generation_mode, init_llm_checkbox, and dcw_enabled. + """ + cfg = get_ui_control_config(is_turbo, is_pure_base=is_pure_base, is_sft=is_sft) + new_choices = cfg["generation_mode_choices"] + if current_mode and current_mode in new_choices: + mode_update = gr.update(choices=new_choices, value=current_mode) + else: + mode_update = gr.update(choices=new_choices) + init_llm_update = gr.update(value=False) if is_pure_base else gr.update() + return ( + gr.update( + value=cfg["inference_steps_value"], + maximum=cfg["inference_steps_maximum"], + minimum=cfg["inference_steps_minimum"], + ), + gr.update(visible=cfg["guidance_scale_visible"]), + gr.update(visible=cfg["use_adg_visible"]), + gr.update(value=cfg["shift_value"], visible=cfg["shift_visible"]), + gr.update(visible=cfg["cfg_interval_start_visible"]), + gr.update(visible=cfg["cfg_interval_end_visible"]), + gr.skip(), # task_type (gr.State — no-op on model config change) + mode_update, + init_llm_update, + gr.update(value=cfg["dcw_enabled_value"]), + ) + + +def get_generation_mode_choices(is_pure_base: bool = False) -> list: + """Get the list of generation mode choices based on model type. + + Args: + is_pure_base: Whether the model is a pure base model. + + Returns: + List of mode choice strings. + """ + if is_pure_base: + return GENERATION_MODES_BASE + else: + return GENERATION_MODES_TURBO diff --git a/experiments/acestep15-xl-sft/overrides/user_preferences.js b/experiments/acestep15-xl-sft/overrides/user_preferences.js new file mode 100644 index 0000000..d134752 --- /dev/null +++ b/experiments/acestep15-xl-sft/overrides/user_preferences.js @@ -0,0 +1,164 @@ +/** + * User preferences persistence – SAVE side only. + * + * Listens for user changes on Gradio UI controls and persists the current + * values to browser localStorage. Restoration is handled on the Python side + * via ``gr.Blocks.load()`` so Gradio's own Svelte reactivity updates every + * component correctly. + * + * Storage schema: + * key = "acestep.ui.user_preferences" + * value = JSON { _version: 2, audio_format: "flac", … } + */ +(() => { + const STORAGE_KEY = "acestep.ui.user_preferences"; + const SCHEMA_VERSION = 2; + const DEBOUNCE_MS = 500; + + /** + * Map of preference key → { elemId, type }. + * elemId : the HTML elem_id set in Gradio + * type : "dropdown" | "slider" | "checkbox" | "number" + */ + const PREFS = { + audio_format: { elemId: "acestep-audio-format", type: "dropdown" }, + mp3_bitrate: { elemId: "acestep-mp3-bitrate", type: "dropdown" }, + mp3_sample_rate: { elemId: "acestep-mp3-sample-rate", type: "dropdown" }, + score_scale: { elemId: "acestep-score-scale", type: "slider" }, + enable_normalization:{ elemId: "acestep-enable-normalization", type: "checkbox" }, + normalization_db: { elemId: "acestep-normalization-db", type: "slider" }, + fade_in_duration: { elemId: "acestep-fade-in-duration", type: "slider" }, + fade_out_duration: { elemId: "acestep-fade-out-duration", type: "slider" }, + latent_shift: { elemId: "acestep-latent-shift", type: "slider" }, + latent_rescale: { elemId: "acestep-latent-rescale", type: "slider" }, + lm_batch_chunk_size: { elemId: "acestep-lm-batch-chunk-size", type: "number" }, + }; + + let saveTimer = null; + const wiredElements = new WeakSet(); + + // ── Storage helpers ────────────────────────────────────────────── + + const saveAll = (prefs) => { + try { + window.localStorage.setItem(STORAGE_KEY, JSON.stringify(prefs)); + } catch (_e) { + // Private browsing or quota exceeded – silently ignore. + } + }; + + // ── DOM helpers ────────────────────────────────────────────────── + + const findInput = (elemId, type) => { + const wrapper = document.getElementById(elemId); + if (!wrapper) return null; + + if (type === "dropdown") { + return wrapper.querySelector("input"); + } + if (type === "slider") { + return wrapper.querySelector("input[type='range']") + || wrapper.querySelector("input[type='number']"); + } + if (type === "checkbox") { + return wrapper.querySelector("input[type='checkbox']"); + } + if (type === "number") { + return wrapper.querySelector("input[type='number']"); + } + return null; + }; + + const readValue = (key) => { + const spec = PREFS[key]; + if (!spec) return undefined; + const el = findInput(spec.elemId, spec.type); + if (!el) return undefined; + + if (spec.type === "checkbox") return el.checked; + if (spec.type === "slider" || spec.type === "number") { + const v = Number(el.value); + return Number.isFinite(v) ? v : undefined; + } + return el.value || undefined; + }; + + // ── Save (debounced) ───────────────────────────────────────────── + + const scheduleSave = () => { + if (saveTimer !== null) { + clearTimeout(saveTimer); + } + saveTimer = setTimeout(() => { + saveTimer = null; + const prefs = { _version: SCHEMA_VERSION }; + for (const key of Object.keys(PREFS)) { + const v = readValue(key); + if (v !== undefined) { + prefs[key] = v; + } + } + saveAll(prefs); + }, DEBOUNCE_MS); + }; + + // ── Wire listeners (re-entrant – safe to call on re-renders) ───── + + const wireListeners = () => { + for (const key of Object.keys(PREFS)) { + const spec = PREFS[key]; + const el = findInput(spec.elemId, spec.type); + if (!el || wiredElements.has(el)) continue; + wiredElements.add(el); + el.addEventListener("input", scheduleSave, { passive: true }); + el.addEventListener("change", scheduleSave, { passive: true }); + } + }; + + // ── MutationObserver – re-wire after Gradio re-renders ─────────── + + const startObserver = () => { + const target = document.getElementById("acestep-audio-format") + || document.body; + const root = target.closest(".gradio-container") || document.body; + + let rafPending = false; + new MutationObserver(() => { + if (rafPending) return; + rafPending = true; + requestAnimationFrame(() => { + rafPending = false; + wireListeners(); + }); + }).observe(root, { childList: true, subtree: true }); + }; + + // ── Boot ───────────────────────────────────────────────────────── + + const BOOT_POLL_MS = 200; + const BOOT_TIMEOUT_MS = 10000; + + const boot = () => { + const started = Date.now(); + const poll = () => { + const probe = document.getElementById( + PREFS.audio_format.elemId + ); + if (!probe) { + if (Date.now() - started < BOOT_TIMEOUT_MS) { + setTimeout(poll, BOOT_POLL_MS); + } + return; + } + wireListeners(); + startObserver(); + }; + poll(); + }; + + if (document.readyState === "loading") { + document.addEventListener("DOMContentLoaded", boot, { once: true }); + } else { + boot(); + } +})(); diff --git a/experiments/acestep15-xl-sft/overrides/user_preferences.py b/experiments/acestep15-xl-sft/overrides/user_preferences.py new file mode 100644 index 0000000..956660c --- /dev/null +++ b/experiments/acestep15-xl-sft/overrides/user_preferences.py @@ -0,0 +1,258 @@ +"""Frontend user-preference persistence helpers for the Gradio UI. + +Save side: A ``" + + +# ── Restore-side: Gradio .load() wiring ───────────────────────────────── + + +def _build_restore_js(num_outputs: int) -> str: + """Build the client-side JS that reads localStorage and returns values. + + The returned function is passed as the ``js`` parameter to + ``demo.load()``. It returns an array whose element order matches + ``PREF_KEYS`` (and therefore the *outputs* list). + + When localStorage has no saved preferences (first visit, cleared + storage, private browsing), the function returns an array of ``null`` + sentinels so the Python side can skip the update and preserve whatever + values were already rendered from ``init_params``. + + Args: + num_outputs: Total number of output components (preference keys + plus any extra outputs like ``mp3_controls_row``). + """ + keys_json = json.dumps(PREF_KEYS) + # Build a type map so the restore JS can validate each value. + type_map: dict[str, str] = {} + for k in PREF_KEYS: + v = _DEFAULTS[k] + if isinstance(v, bool): + type_map[k] = "boolean" + elif isinstance(v, (int, float)): + type_map[k] = "number" + else: + type_map[k] = "string" + type_map_json = json.dumps(type_map, ensure_ascii=False) + # Keys whose Gradio Dropdown choices are integers stored as strings in + # localStorage. Only actual dropdown keys with numeric defaults need + # coercion; sliders/numbers are already stored as numbers. + numeric_dropdown_keys_json = json.dumps(["mp3_sample_rate"]) + # Sentinel array returned when there is nothing to restore. Using null + # lets the Python fn detect "no stored prefs" and return gr.update() + # for every output, preserving the values already rendered on the page. + skip_sentinel = f"new Array({num_outputs}).fill(null)" + return f"""() => {{ + const STORAGE_KEY = {json.dumps(_STORAGE_KEY)}; + const SCHEMA_VERSION = {_SCHEMA_VERSION}; + const KEYS = {keys_json}; + const TYPE_MAP = {type_map_json}; + const NUMERIC_COERCE_KEYS = new Set({numeric_dropdown_keys_json}); + const SKIP = {skip_sentinel}; + try {{ + const raw = window.localStorage.getItem(STORAGE_KEY); + if (!raw) return SKIP; + const prefs = JSON.parse(raw); + // Only reset on downgrade; forward-compatible additions of new + // keys are handled by skipping (preserving init_params). + if (prefs._version !== SCHEMA_VERSION) {{ + return SKIP; + }} + const result = KEYS.map(k => {{ + if (!(k in prefs)) return null; + let v = prefs[k]; + // Type-check: fall back to null (skip) if the stored type + // does not match what the Gradio component expects. + const expected = TYPE_MAP[k]; + if (expected && typeof v !== expected) {{ + // Allow stringified numbers for dropdown coercion below. + if (!(NUMERIC_COERCE_KEYS.has(k) && typeof v === "string")) {{ + return null; + }} + }} + // Coerce stringified numbers back for Dropdown choices that + // expect integers (e.g. mp3_sample_rate: 48000 not "48000"). + if (NUMERIC_COERCE_KEYS.has(k) && typeof v === "string") {{ + const n = Number(v); + if (Number.isFinite(n)) v = n; + else return null; + }} + return v; + }}); + // If none of the keys had stored values, skip entirely. + if (result.every(v => v === null)) return SKIP; + // Compute mp3 control visibility from audio_format (index 0). + // Push 3 extra values: mp3_controls_row, mp3_bitrate, mp3_sample_rate + // matching the outputs of _update_mp3_control_visibility(). + // When audioFormat is null (no stored value), push nulls so Python + // emits gr.update() and preserves whatever init_params set. + const audioFormat = result[0]; + const mp3 = audioFormat === null ? null : audioFormat === "mp3"; + result.push(mp3, mp3, mp3); + return result; + }} catch (_e) {{ + return SKIP; + }} + }}""" + + +def restore_preferences( + *values: Any, _num_outputs: int = 0 +) -> tuple[Any, ...]: + """Map JS restore results into Gradio output values. + + The JS function reads localStorage and produces an array: + - First ``len(PREF_KEYS)`` elements are preference values (or null). + - Next 3 elements are mp3 visibility booleans (or null): + [mp3_controls_row, mp3_bitrate, mp3_sample_rate]. + + ``None`` (JSON ``null``) → ``gr.update()`` (no-op, preserves current). + Booleans beyond PREF_KEYS → visibility/interactivity updates matching + ``_update_mp3_control_visibility()`` from the output controls module. + + When the JS side returns no values (e.g. certain Gradio versions do not + forward the JS return value to the Python ``fn`` when ``inputs=None``), + ``_num_outputs`` is used to produce the correct number of no-op updates + so Gradio does not raise a ``ValueError`` about mismatched output count. + """ + import gradio as gr + + if not values: + return tuple(gr.update() for _ in range(_num_outputs)) + + n_prefs = len(PREF_KEYS) + results: list[Any] = [] + for i, v in enumerate(values): + if v is None: + results.append(gr.update()) + elif i == n_prefs and isinstance(v, bool): + # mp3_controls_row: visibility only. + results.append(gr.update(visible=v)) + elif i > n_prefs and isinstance(v, bool): + # mp3_bitrate, mp3_sample_rate: visibility + interactivity. + results.append(gr.update(visible=v, interactive=v)) + else: + results.append(v) + return tuple(results) + + +def wire_preference_restore( + demo: Any, + generation_section: dict[str, Any], + *, + service_mode: bool = False, +) -> None: + """Attach a ``demo.load()`` handler that restores saved preferences. + + Must be called **inside** the ``with gr.Blocks() as demo:`` context, + after all generation components have been created. + + In service mode the function is a no-op: service-mode sessions use + server-side ``init_params`` and controls are locked + (``interactive=False``), so localStorage values must not override them. + + Args: + demo: The ``gr.Blocks`` instance. + generation_section: Merged component dict that includes the output + control components (``audio_format``, ``mp3_bitrate``, etc.). + service_mode: When ``True``, skip wiring entirely so that + localStorage cannot override server-configured values. + """ + if service_mode: + return + + outputs = [] + for key in PREF_KEYS: + component = generation_section.get(key) + if component is None: + raise KeyError( + f"wire_preference_restore: missing component {key!r} in " + f"generation_section (available: {sorted(generation_section)})" + ) + outputs.append(component) + + # Also update mp3 control visibility so it stays in sync when the + # restored audio_format differs from the server-rendered default. + # Gradio does not fire .change() for load-time value assignments, so + # without this the MP3 row and its children could be visible/hidden + # incorrectly. The three extra outputs mirror the return of + # _update_mp3_control_visibility(): [row, bitrate, sample_rate]. + for mp3_key in ("mp3_controls_row", "mp3_bitrate", "mp3_sample_rate"): + comp = generation_section.get(mp3_key) + if comp is not None: + outputs.append(comp) + + demo.load( + fn=partial(restore_preferences, _num_outputs=len(outputs)), + inputs=None, + outputs=outputs, + js=_build_restore_js(num_outputs=len(outputs)), + )