Files

138 lines
7.7 KiB
Python

"""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")