mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-30 17:05:27 +08:00
fix(webui): crash on non-ffmpeg env. (#466)
This commit is contained in:
+3
-9
@@ -12,23 +12,17 @@ from io import BytesIO
|
||||
|
||||
import ChatTTS
|
||||
|
||||
from tools.audio import unsafe_float_to_int16, wav2
|
||||
from tools.audio import wav_arr_to_mp3_view
|
||||
from tools.logger import get_logger
|
||||
|
||||
logger = get_logger("Command")
|
||||
|
||||
|
||||
def save_mp3_file(wav, index):
|
||||
buf = BytesIO()
|
||||
with wave.open(buf, "wb") as wf:
|
||||
wf.setnchannels(1) # Mono channel
|
||||
wf.setsampwidth(2) # Sample width in bytes
|
||||
wf.setframerate(24000) # Sample rate in Hz
|
||||
wf.writeframes(unsafe_float_to_int16(wav))
|
||||
buf.seek(0, 0)
|
||||
data = wav_arr_to_mp3_view(wav)
|
||||
mp3_filename = f"output_audio_{index}.mp3"
|
||||
with open(mp3_filename, "wb") as f:
|
||||
wav2(buf, f, "mp3")
|
||||
f.write(data)
|
||||
logger.info(f"Audio saved to {mp3_filename}")
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from time import sleep
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
|
||||
from tools.audio import unsafe_float_to_int16
|
||||
from tools.audio import wav_arr_to_mp3_view
|
||||
from tools.logger import get_logger
|
||||
|
||||
logger = get_logger(" WebUI ")
|
||||
@@ -146,10 +146,10 @@ def generate_audio(text, temperature, top_P, top_K, audio_seed_input, 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])
|
||||
yield wav_arr_to_mp3_view(audio[0]).tobytes()
|
||||
del audio
|
||||
else:
|
||||
yield 24000, unsafe_float_to_int16(np.array(wav[0]).flatten())
|
||||
yield wav_arr_to_mp3_view(np.array(wav[0]).flatten()).tobytes()
|
||||
|
||||
|
||||
def interrupt_generate():
|
||||
|
||||
@@ -107,7 +107,6 @@ def main():
|
||||
streaming=stream,
|
||||
interactive=False,
|
||||
show_label=True,
|
||||
format="mp3",
|
||||
)
|
||||
generate_button.click(fn=set_buttons_before_generate, inputs=[generate_button, interrupt_button], outputs=[generate_button, interrupt_button]).then(
|
||||
refine_text,
|
||||
|
||||
@@ -1,2 +1 @@
|
||||
from .np import unsafe_float_to_int16
|
||||
from .av import wav2
|
||||
from .mp3 import wav_arr_to_mp3_view
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
import wave
|
||||
from io import BytesIO
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .np import unsafe_float_to_int16
|
||||
from .av import wav2
|
||||
|
||||
def wav_arr_to_mp3_view(wav: np.ndarray):
|
||||
buf = BytesIO()
|
||||
with wave.open(buf, "wb") as wf:
|
||||
wf.setnchannels(1) # Mono channel
|
||||
wf.setsampwidth(2) # Sample width in bytes
|
||||
wf.setframerate(24000) # Sample rate in Hz
|
||||
wf.writeframes(unsafe_float_to_int16(wav))
|
||||
buf.seek(0, 0)
|
||||
buf2 = BytesIO()
|
||||
wav2(buf, buf2, "mp3")
|
||||
buf.seek(0, 0)
|
||||
return buf2.getbuffer()
|
||||
Reference in New Issue
Block a user