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

@@ -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