ui components support shift change and cfg type res multistep

This commit is contained in:
Gong Junmin
2025-05-08 17:02:47 +08:00
parent 2097aca44b
commit cf227c9c37
2 changed files with 20 additions and 1 deletions
+3
View File
@@ -1396,6 +1396,7 @@ class ACEStepPipeline:
use_erg_tag: bool = True, use_erg_tag: bool = True,
use_erg_lyric: bool = True, use_erg_lyric: bool = True,
use_erg_diffusion: bool = True, use_erg_diffusion: bool = True,
shift: float = 3.0,
oss_steps: str = None, oss_steps: str = None,
guidance_scale_text: float = 0.0, guidance_scale_text: float = 0.0,
guidance_scale_lyric: float = 0.0, guidance_scale_lyric: float = 0.0,
@@ -1586,6 +1587,7 @@ class ACEStepPipeline:
repaint_start=repaint_start, repaint_start=repaint_start,
repaint_end=repaint_end, repaint_end=repaint_end,
src_latents=src_latents, src_latents=src_latents,
shift=shift,
) )
end_time = time.time() end_time = time.time()
@@ -1638,6 +1640,7 @@ class ACEStepPipeline:
"src_audio_path": src_audio_path, "src_audio_path": src_audio_path,
"edit_target_prompt": edit_target_prompt, "edit_target_prompt": edit_target_prompt,
"edit_target_lyrics": edit_target_lyrics, "edit_target_lyrics": edit_target_lyrics,
"shift": shift,
} }
# save input_params_json # save input_params_json
for output_audio_path in output_paths: for output_audio_path in output_paths:
+17 -1
View File
@@ -146,7 +146,7 @@ def create_text2music_ui(
with gr.Accordion("Advanced Settings", open=False): with gr.Accordion("Advanced Settings", open=False):
scheduler_type = gr.Radio( scheduler_type = gr.Radio(
["euler", "heun"], ["euler", "heun", "res_multistep"],
value="euler", value="euler",
label="Scheduler Type", label="Scheduler Type",
elem_id="scheduler_type", elem_id="scheduler_type",
@@ -218,6 +218,14 @@ def create_text2music_ui(
value=None, value=None,
info="Optimal Steps for the generation. But not test well", info="Optimal Steps for the generation. But not test well",
) )
shift = gr.Slider(
minimum=1.0,
maximum=5.0,
step=0.1,
value=3.0,
label="shift",
interactive=True,
)
text2music_bnt = gr.Button("Generate", variant="primary") text2music_bnt = gr.Button("Generate", variant="primary")
@@ -263,6 +271,7 @@ def create_text2music_ui(
), ),
retake_seeds=retake_seeds, retake_seeds=retake_seeds,
retake_variance=retake_variance, retake_variance=retake_variance,
shift=shift,
task="retake", task="retake",
) )
@@ -381,6 +390,7 @@ def create_text2music_ui(
guidance_scale_lyric, guidance_scale_lyric,
retake_seeds=retake_seeds, retake_seeds=retake_seeds,
retake_variance=retake_variance, retake_variance=retake_variance,
shift=shift,
task="repaint", task="repaint",
repaint_start=repaint_start, repaint_start=repaint_start,
repaint_end=repaint_end, repaint_end=repaint_end,
@@ -696,6 +706,7 @@ def create_text2music_ui(
guidance_scale_lyric, guidance_scale_lyric,
retake_seeds=extend_seeds, retake_seeds=extend_seeds,
retake_variance=1.0, retake_variance=1.0,
shift=shift,
task="extend", task="extend",
repaint_start=repaint_start, repaint_start=repaint_start,
repaint_end=repaint_end, repaint_end=repaint_end,
@@ -751,6 +762,11 @@ def create_text2music_ui(
json_data["use_erg_tag"], json_data["use_erg_tag"],
json_data["use_erg_lyric"], json_data["use_erg_lyric"],
json_data["use_erg_diffusion"], json_data["use_erg_diffusion"],
(
json_data["shift"]
if "shift" in json_data
else 3.0
),
", ".join(map(str, json_data["oss_steps"])), ", ".join(map(str, json_data["oss_steps"])),
( (
json_data["guidance_scale_text"] json_data["guidance_scale_text"]