Persist ACE-Step XL-SFT quality defaults
This commit is contained in:
1 parent
eeebbd06eb
commit
56c382f71f
6 files changed
+868
No files matched your search
@@ -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)),
|
||||||
|
)
|
||||||
Reference in new issue
Block a user