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
@@ -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
|
||||
Reference in new issue
Block a user