feat: interruptible infer (#433)

This commit is contained in:
源文雨
2024-06-25 01:33:07 +09:00
committed by GitHub
parent 5222976ed2
commit c23e514b0a
5 changed files with 77 additions and 23 deletions
+41 -16
View File
@@ -17,6 +17,8 @@ chat = ChatTTS.Chat(get_logger("ChatTTS"))
custom_path: Optional[str] = None
has_interrupted = False
# 音色选项:用于预置合适的音色
voices = {
"Default": {"seed": 2},
@@ -83,24 +85,31 @@ def reload_chat(coef: Optional[str]) -> str:
gr.Info("Reload succeess.")
return chat.coef
def set_generate_buttons(generate_button, interrupt_button, is_reset=False):
return gr.update(value=generate_button, visible=is_reset, interactive=is_reset), gr.update(value=interrupt_button, visible=not is_reset, interactive=not is_reset)
def refine_text(text, text_seed_input, refine_text_flag, generate_button, interrupt_button):
global chat, has_interrupted
has_interrupted = False
def refine_text(text, text_seed_input, refine_text_flag):
if not refine_text_flag:
return text
global chat
return text, *set_generate_buttons(generate_button, interrupt_button, is_reset=True)
with TorchSeedContext(text_seed_input):
text = chat.infer(text,
skip_refine_text=False,
refine_text_only=True,
)
return text[0] if isinstance(text, list) else text
text = chat.infer(
text,
skip_refine_text=False,
refine_text_only=True,
)
return text[0] if isinstance(text, list) else text, *set_generate_buttons(generate_button, interrupt_button, is_reset=True)
def text_output_listener(generate_button, interrupt_button):
return set_generate_buttons(generate_button, interrupt_button)
def generate_audio(text, temperature, top_P, top_K, audio_seed_input, stream):
if not text: return None
global chat, has_interrupted
global chat
if not text or text == "𝕃𝕠𝕒𝕕𝕚𝕟𝕘..." or has_interrupted: return None
with TorchSeedContext(audio_seed_input):
rand_spk = chat.sample_random_speaker()
@@ -119,10 +128,26 @@ def generate_audio(text, temperature, top_P, top_K, audio_seed_input, stream):
params_infer_code=params_infer_code,
stream=stream,
)
if stream:
for gen in wav:
yield 24000, unsafe_float_to_int16(gen[0][0])
return
if stream:
for gen in wav:
audio = gen[0]
if audio is not None and len(audio) > 0:
yield 24000, unsafe_float_to_int16(audio[0])
del audio
return
yield 24000, unsafe_float_to_int16(np.array(wav[0]).flatten())
def interrupt_generate():
global chat, has_interrupted
has_interrupted = True
chat.interrupt()
def set_buttons_after_generate(generate_button, interrupt_button, audio_output):
global has_interrupted
return set_generate_buttons(
generate_button, interrupt_button,
audio_output is not None or has_interrupted,
)
+8 -4
View File
@@ -48,6 +48,7 @@ def main():
auto_play_checkbox = gr.Checkbox(label="Auto Play", value=False, scale=1)
stream_mode_checkbox = gr.Checkbox(label="Stream Mode", value=False, scale=1)
generate_button = gr.Button("Generate", scale=2, variant="primary")
interrupt_button = gr.Button("Interrupt", scale=2, variant="stop", visible=False, interactive=False)
text_output = gr.Textbox(label="Output Text", interactive=False)
@@ -64,10 +65,12 @@ def main():
reload_chat_button.click(reload_chat, inputs=dvae_coef_text, outputs=dvae_coef_text)
generate_button.click(fn=lambda: "", outputs=text_output)
generate_button.click(fn=lambda: "𝕃𝕠𝕒𝕕𝕚𝕟𝕘...", outputs=text_output)
generate_button.click(refine_text,
inputs=[text_input, text_seed_input, refine_text_checkbox],
outputs=text_output)
inputs=[text_input, text_seed_input, refine_text_checkbox, generate_button, interrupt_button],
outputs=[text_output, generate_button, interrupt_button])
interrupt_button.click(interrupt_generate)
@gr.render(inputs=[auto_play_checkbox, stream_mode_checkbox])
def make_audio(autoplay, stream):
@@ -79,9 +82,10 @@ def main():
interactive=False,
show_label=True,
)
text_output.change(text_output_listener, inputs=[generate_button, interrupt_button], outputs=[generate_button, interrupt_button])
text_output.change(generate_audio,
inputs=[text_output, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, stream_mode_checkbox],
outputs=audio_output)
outputs=audio_output).then(fn=set_buttons_after_generate, inputs=[generate_button, interrupt_button, audio_output], outputs=[generate_button, interrupt_button])
gr.Examples(
examples=[