Persist ACE-Step XL-SFT quality defaults
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)),
|
||||
)
|
||||
Reference in New Issue
Block a user