Persist ACE-Step XL-SFT quality defaults

This commit is contained in:
Mikei386 committed 2026-09-08 16:37:36 +02:00
1 parent eeebbd06eb
commit 56c382f71f
6 files changed
+868

No files matched your search

+21
View File
@@ -9,6 +9,27 @@ Das offizielle Image ist auf den am 8. September 2026 geladenen Digest
`sha256:95652cd780c78a1b1a7f6f0335530430f0ae53d96c7c12d59f9f39fa23d38567` `sha256:95652cd780c78a1b1a7f6f0335530430f0ae53d96c7c12d59f9f39fa23d38567`
fixiert. 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 ## Start
Vor dem Start muessen das aktive llama.cpp-Profil und Qwen3-TTS beendet sein. Vor dem Start muessen das aktive llama.cpp-Profil und Qwen3-TTS beendet sein.
@@ -11,6 +11,9 @@ services:
ACESTEP_INIT_SERVICE: "true" ACESTEP_INIT_SERVICE: "true"
ACESTEP_INIT_LLM: "true" ACESTEP_INIT_LLM: "true"
ACESTEP_DEVICE: cuda 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" TOKENIZERS_PARALLELISM: "false"
NVIDIA_VISIBLE_DEVICES: ${ACESTEP_GPU_UUID:?set ACESTEP_GPU_UUID to the RTX 5080 UUID} NVIDIA_VISIBLE_DEVICES: ${ACESTEP_GPU_UUID:?set ACESTEP_GPU_UUID to the RTX 5080 UUID}
deploy: deploy:
@@ -27,6 +30,11 @@ services:
- ${ACESTEP_HF_CACHE_DIR:-/data/models/acestep/hf-cache}:/root/.cache/huggingface - ${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/gradio_outputs
- ${ACESTEP_OUTPUT_DIR:-/data/music/acestep}:/app/output - ${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" shm_size: "2gb"
restart: "no" restart: "no"
healthcheck: healthcheck:
@@ -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}
@@ -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
@@ -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();
}
})();
@@ -0,0 +1,258 @@
"""Frontend user-preference persistence helpers for the Gradio UI.
Save side: A ``<script>`` injected via ``Blocks(head=…)`` listens for DOM
changes and writes the current preference values to ``localStorage``.
Restore side: ``wire_preference_restore`` attaches a ``demo.load()`` handler
whose *js* parameter reads ``localStorage`` on page load and feeds the saved
values straight into the Gradio component outputs. Because Gradio itself
applies the updates through its own Svelte reactivity, every component type
(dropdown, slider, checkbox, number) is updated correctly—no fragile DOM
hacking required.
"""
from __future__ import annotations
import json
from functools import partial
from pathlib import Path
from typing import Any
_ASSET_FILENAME = "user_preferences.js"
_STORAGE_KEY = "acestep.ui.user_preferences"
_SCHEMA_VERSION = 2
# Ordered list of preference keys. The order here MUST match the order of
# *outputs* passed to ``demo.load()`` in ``wire_preference_restore``.
PREF_KEYS: list[str] = [
"audio_format",
"mp3_bitrate",
"mp3_sample_rate",
"score_scale",
"enable_normalization",
"normalization_db",
"fade_in_duration",
"fade_out_duration",
"latent_shift",
"latent_rescale",
"lm_batch_chunk_size",
]
# Default values used when localStorage is empty or the schema version has
# changed. Keys must match ``PREF_KEYS``.
_DEFAULTS: dict[str, Any] = {
"audio_format": "flac",
"mp3_bitrate": "320k",
"mp3_sample_rate": 48000,
"score_scale": 0.5,
"enable_normalization": True,
"normalization_db": -1.0,
"fade_in_duration": 0.0,
"fade_out_duration": 0.0,
"latent_shift": 0.0,
"latent_rescale": 1.0,
"lm_batch_chunk_size": 8,
}
# ── Save-side: head script injection ────────────────────────────────────
def _load_preferences_script() -> str:
"""Load the external save-preferences JavaScript asset."""
asset_path = Path(__file__).with_name(_ASSET_FILENAME)
return asset_path.read_text(encoding="utf-8").strip()
def get_user_preferences_head() -> str:
"""Return Gradio head HTML that injects save-side preference persistence."""
script_source = _load_preferences_script()
return f"<script>\n{script_source}\n</script>"
# ── 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)),
)