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