Fix ACE-Step community UI parameter wiring
This commit is contained in:
@@ -0,0 +1,137 @@
|
||||
"""Extend ACE-Step's official /release_task route with named generation inputs.
|
||||
|
||||
The base image already provides the route. This build-time patch only exposes
|
||||
the parameters supported by its installed GenerationParams/GenerationConfig
|
||||
dataclasses, so the separate Community UI never has to depend on Gradio's
|
||||
positional component order.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
|
||||
target = Path(sys.argv[1])
|
||||
source = target.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def replace_once(old: str, new: str, label: str) -> None:
|
||||
global source
|
||||
count = source.count(old)
|
||||
if count != 1:
|
||||
raise RuntimeError(f"{label}: expected one anchor, found {count}")
|
||||
source = source.replace(old, new, 1)
|
||||
|
||||
|
||||
old_params = ''' # Build generation params with alias support
|
||||
params = GenerationParams(
|
||||
task_type=get_param("task_type", default="text2music"),
|
||||
caption=caption,
|
||||
lyrics=lyrics,
|
||||
bpm=sample_bpm or get_param("bpm"),
|
||||
keyscale=sample_keyscale or get_param("key_scale", "keyscale", "key", default=""),
|
||||
timesignature=sample_timesignature or get_param("time_signature", "timesignature", default=""),
|
||||
duration=sample_duration or get_param("audio_duration", "duration", default=-1),
|
||||
vocal_language=sample_language,
|
||||
inference_steps=get_param("inference_steps", default=8),
|
||||
guidance_scale=float(get_param("guidance_scale", default=7.0) or 7.0),
|
||||
seed=int(get_param("seed", default=-1) or -1),
|
||||
thinking=to_bool(get_param("thinking"), False),
|
||||
lm_temperature=lm_temperature,
|
||||
lm_cfg_scale=float(get_param("lm_cfg_scale", default=2.0) or 2.0),
|
||||
lm_negative_prompt=get_param("lm_negative_prompt", default="NO USER INPUT") or "NO USER INPUT",
|
||||
repaint_latent_crossfade_frames=int(
|
||||
get_param("repaint_latent_crossfade_frames", default=10) or 10,
|
||||
),
|
||||
repaint_wav_crossfade_sec=float(
|
||||
get_param("repaint_wav_crossfade_sec", default=0.0) or 0.0,
|
||||
),
|
||||
repaint_mode=get_param("repaint_mode", default="balanced") or "balanced",
|
||||
repaint_strength=float(
|
||||
get_param("repaint_strength", default=0.5) or 0.5,
|
||||
),
|
||||
)
|
||||
'''
|
||||
|
||||
new_params = ''' # Build generation params with alias support. Keep this
|
||||
# mapping explicit: every public API field below is named and independent
|
||||
# from the order of components in the Gradio interface.
|
||||
raw_bpm = sample_bpm or get_param("bpm")
|
||||
params = GenerationParams(
|
||||
task_type=get_param("task_type", default="text2music") or "text2music",
|
||||
instruction=get_param("instruction", default="Fill the audio semantic mask based on the given conditions:") or "Fill the audio semantic mask based on the given conditions:",
|
||||
reference_audio=get_param("reference_audio_path", "reference_audio"),
|
||||
src_audio=get_param("src_audio_path", "src_audio", "source_audio"),
|
||||
audio_codes=get_param("audio_codes", default="") or "",
|
||||
caption=caption,
|
||||
lyrics=lyrics,
|
||||
instrumental=to_bool(get_param("instrumental"), False),
|
||||
bpm=int(float(raw_bpm)) if raw_bpm not in (None, "", 0, "0") else None,
|
||||
keyscale=sample_keyscale or get_param("key_scale", "keyscale", "key", default=""),
|
||||
timesignature=sample_timesignature or get_param("time_signature", "timesignature", default=""),
|
||||
duration=float(sample_duration or get_param("audio_duration", "duration", default=-1) or -1),
|
||||
vocal_language=sample_language,
|
||||
inference_steps=int(get_param("inference_steps", default=50) or 50),
|
||||
guidance_scale=float(get_param("guidance_scale", default=7.0) or 7.0),
|
||||
seed=int(get_param("seed", default=-1) or -1),
|
||||
use_adg=to_bool(get_param("use_adg"), False),
|
||||
cfg_interval_start=float(get_param("cfg_interval_start", default=0.0) or 0.0),
|
||||
cfg_interval_end=float(get_param("cfg_interval_end", default=1.0) or 1.0),
|
||||
shift=float(get_param("shift", default=1.0) or 1.0),
|
||||
infer_method=get_param("infer_method", default="ode") or "ode",
|
||||
sampler_mode=get_param("sampler_mode", default="euler") or "euler",
|
||||
repainting_start=float(get_param("repainting_start", default=0.0) or 0.0),
|
||||
repainting_end=float(get_param("repainting_end", default=-1.0) or -1.0),
|
||||
chunk_mask_mode=get_param("chunk_mask_mode", default="auto") or "auto",
|
||||
audio_cover_strength=float(get_param("audio_cover_strength", default=1.0) or 1.0),
|
||||
cover_noise_strength=float(get_param("cover_noise_strength", default=0.0) or 0.0),
|
||||
thinking=to_bool(get_param("thinking"), True),
|
||||
lm_temperature=lm_temperature,
|
||||
lm_cfg_scale=float(get_param("lm_cfg_scale", default=2.0) or 2.0),
|
||||
lm_top_k=int(get_param("lm_top_k", default=0) or 0),
|
||||
lm_top_p=float(get_param("lm_top_p", default=0.9) or 0.9),
|
||||
lm_negative_prompt=get_param("lm_negative_prompt", default="NO USER INPUT") or "NO USER INPUT",
|
||||
use_cot_metas=to_bool(get_param("use_cot_metas"), True),
|
||||
use_cot_caption=to_bool(get_param("use_cot_caption"), True),
|
||||
use_cot_lyrics=to_bool(get_param("use_cot_lyrics"), False),
|
||||
use_cot_language=to_bool(get_param("use_cot_language"), True),
|
||||
use_constrained_decoding=to_bool(get_param("use_constrained_decoding"), True),
|
||||
enable_normalization=to_bool(get_param("enable_normalization"), True),
|
||||
normalization_db=float(get_param("normalization_db", default=-1.0) or -1.0),
|
||||
fade_in_duration=float(get_param("fade_in_duration", default=0.0) or 0.0),
|
||||
fade_out_duration=float(get_param("fade_out_duration", default=0.0) or 0.0),
|
||||
latent_shift=float(get_param("latent_shift", default=0.0) or 0.0),
|
||||
latent_rescale=float(get_param("latent_rescale", default=1.0) or 1.0),
|
||||
repaint_latent_crossfade_frames=int(get_param("repaint_latent_crossfade_frames", default=10) or 10),
|
||||
repaint_wav_crossfade_sec=float(get_param("repaint_wav_crossfade_sec", default=0.0) or 0.0),
|
||||
repaint_mode=get_param("repaint_mode", default="balanced") or "balanced",
|
||||
repaint_strength=float(get_param("repaint_strength", default=0.5) or 0.5),
|
||||
)
|
||||
'''
|
||||
|
||||
replace_once(old_params, new_params, "GenerationParams mapping")
|
||||
|
||||
old_config = ''' config = GenerationConfig(
|
||||
batch_size=get_param("batch_size", default=2),
|
||||
use_random_seed=use_random_seed,
|
||||
seeds=resolved_seeds,
|
||||
audio_format=get_param("audio_format", default="flac"),
|
||||
mp3_bitrate=get_param("mp3_bitrate", default="128k"),
|
||||
mp3_sample_rate=get_param("mp3_sample_rate", default=48000),
|
||||
)
|
||||
'''
|
||||
|
||||
new_config = ''' config = GenerationConfig(
|
||||
batch_size=int(get_param("batch_size", default=1) or 1),
|
||||
allow_lm_batch=to_bool(get_param("allow_lm_batch"), True),
|
||||
use_random_seed=to_bool(use_random_seed, True),
|
||||
seeds=resolved_seeds,
|
||||
lm_batch_chunk_size=int(get_param("lm_batch_chunk_size", default=8) or 8),
|
||||
constrained_decoding_debug=to_bool(get_param("constrained_decoding_debug"), False),
|
||||
audio_format=get_param("audio_format", default="flac") or "flac",
|
||||
mp3_bitrate=get_param("mp3_bitrate", default="320k") or "320k",
|
||||
mp3_sample_rate=int(get_param("mp3_sample_rate", default=48000) or 48000),
|
||||
)
|
||||
'''
|
||||
|
||||
replace_once(old_config, new_config, "GenerationConfig mapping")
|
||||
target.write_text(source, encoding="utf-8")
|
||||
Reference in New Issue
Block a user