mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-09-01 14:55:37 +08:00
feat: interruptible infer (#433)
This commit is contained in:
+41
-16
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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=[
|
||||
|
||||
Reference in New Issue
Block a user