mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 17:43:55 +08:00
update:合并main分支
This commit is contained in:
@@ -239,7 +239,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if self._is_token_expired():
|
if self._is_token_expired():
|
||||||
logger.warning("Token已过期,正在自动刷新...")
|
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
|
||||||
self._refresh_token()
|
self._refresh_token()
|
||||||
|
|
||||||
file_path = None
|
file_path = None
|
||||||
|
|||||||
@@ -6,18 +6,18 @@ import hashlib
|
|||||||
import base64
|
import base64
|
||||||
import requests
|
import requests
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
|
|
||||||
from core.providers.tts.base import TTSProviderBase
|
|
||||||
|
|
||||||
import http.client
|
|
||||||
import urllib.parse
|
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from urllib import parse
|
from urllib import parse
|
||||||
|
|
||||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class AccessToken:
|
class AccessToken:
|
||||||
@@ -48,7 +48,7 @@ class AccessToken:
|
|||||||
}
|
}
|
||||||
# 构造规范化的请求字符串
|
# 构造规范化的请求字符串
|
||||||
query_string = AccessToken._encode_dict(parameters)
|
query_string = AccessToken._encode_dict(parameters)
|
||||||
print("规范化的请求字符串: %s" % query_string)
|
# print('规范化的请求字符串: %s' % query_string)
|
||||||
# 构造待签名字符串
|
# 构造待签名字符串
|
||||||
string_to_sign = (
|
string_to_sign = (
|
||||||
"GET"
|
"GET"
|
||||||
@@ -57,7 +57,7 @@ class AccessToken:
|
|||||||
+ "&"
|
+ "&"
|
||||||
+ AccessToken._encode_text(query_string)
|
+ AccessToken._encode_text(query_string)
|
||||||
)
|
)
|
||||||
print("待签名的字符串: %s" % string_to_sign)
|
# print('待签名的字符串: %s' % string_to_sign)
|
||||||
# 计算签名
|
# 计算签名
|
||||||
secreted_string = hmac.new(
|
secreted_string = hmac.new(
|
||||||
bytes(access_key_secret + "&", encoding="utf-8"),
|
bytes(access_key_secret + "&", encoding="utf-8"),
|
||||||
@@ -65,16 +65,16 @@ class AccessToken:
|
|||||||
hashlib.sha1,
|
hashlib.sha1,
|
||||||
).digest()
|
).digest()
|
||||||
signature = base64.b64encode(secreted_string)
|
signature = base64.b64encode(secreted_string)
|
||||||
print("签名: %s" % signature)
|
# print('签名: %s' % signature)
|
||||||
# 进行URL编码
|
# 进行URL编码
|
||||||
signature = AccessToken._encode_text(signature)
|
signature = AccessToken._encode_text(signature)
|
||||||
print("URL编码后的签名: %s" % signature)
|
# print('URL编码后的签名: %s' % signature)
|
||||||
# 调用服务
|
# 调用服务
|
||||||
full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (
|
full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (
|
||||||
signature,
|
signature,
|
||||||
query_string,
|
query_string,
|
||||||
)
|
)
|
||||||
print("url: %s" % full_url)
|
# print('url: %s' % full_url)
|
||||||
# 提交HTTP GET请求
|
# 提交HTTP GET请求
|
||||||
response = requests.get(full_url)
|
response = requests.get(full_url)
|
||||||
if response.ok:
|
if response.ok:
|
||||||
@@ -84,7 +84,7 @@ class AccessToken:
|
|||||||
token = root_obj[key]["Id"]
|
token = root_obj[key]["Id"]
|
||||||
expire_time = root_obj[key]["ExpireTime"]
|
expire_time = root_obj[key]["ExpireTime"]
|
||||||
return token, expire_time
|
return token, expire_time
|
||||||
print(response.text)
|
# print(response.text)
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
@@ -94,33 +94,67 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
|
|
||||||
# 新增空值判断逻辑
|
# 新增空值判断逻辑
|
||||||
access_key_id = config.get("access_key_id")
|
self.access_key_id = config.get("access_key_id")
|
||||||
access_key_secret = config.get("access_key_secret")
|
self.access_key_secret = config.get("access_key_secret")
|
||||||
if access_key_id and access_key_secret:
|
|
||||||
# 使用密钥对生成临时token
|
|
||||||
token, expire_time = AccessToken.create_token(
|
|
||||||
access_key_id, access_key_secret
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# 直接使用预生成的长期token
|
|
||||||
token = config.get("token")
|
|
||||||
expire_time = None
|
|
||||||
|
|
||||||
print("token: %s, expire time(s): %s" % (token, expire_time))
|
|
||||||
|
|
||||||
self.appkey = config.get("appkey")
|
self.appkey = config.get("appkey")
|
||||||
self.token = token
|
|
||||||
self.format = config.get("format", "wav")
|
self.format = config.get("format", "wav")
|
||||||
self.sample_rate = config.get("sample_rate", 16000)
|
self.sample_rate = config.get("sample_rate", 16000)
|
||||||
self.voice = config.get("voice", "xiaoyun")
|
self.voice = config.get("voice", "xiaoyun")
|
||||||
self.volume = config.get("volume", 50)
|
self.volume = config.get("volume", 50)
|
||||||
self.speech_rate = config.get("speech_rate", 0)
|
self.speech_rate = config.get("speech_rate", 0)
|
||||||
self.pitch_rate = config.get("pitch_rate", 0)
|
self.pitch_rate = config.get("pitch_rate", 0)
|
||||||
|
|
||||||
self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com")
|
self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com")
|
||||||
self.api_url = f"https://{self.host}/stream/v1/tts"
|
self.api_url = f"https://{self.host}/stream/v1/tts"
|
||||||
self.header = {"Content-Type": "application/json"}
|
self.header = {"Content-Type": "application/json"}
|
||||||
|
|
||||||
|
if self.access_key_id and self.access_key_secret:
|
||||||
|
# 使用密钥对生成临时token
|
||||||
|
self._refresh_token()
|
||||||
|
else:
|
||||||
|
# 直接使用预生成的长期token
|
||||||
|
self.token = config.get("token")
|
||||||
|
self.expire_time = None
|
||||||
|
|
||||||
|
def _refresh_token(self):
|
||||||
|
"""刷新Token并记录过期时间"""
|
||||||
|
if self.access_key_id and self.access_key_secret:
|
||||||
|
self.token, expire_time_str = AccessToken.create_token(
|
||||||
|
self.access_key_id, self.access_key_secret
|
||||||
|
)
|
||||||
|
if not expire_time_str:
|
||||||
|
raise ValueError("无法获取有效的Token过期时间")
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 统一转换为字符串处理
|
||||||
|
expire_str = str(expire_time_str).strip()
|
||||||
|
|
||||||
|
if expire_str.isdigit():
|
||||||
|
expire_time = datetime.fromtimestamp(int(expire_str))
|
||||||
|
else:
|
||||||
|
expire_time = datetime.strptime(expire_str, "%Y-%m-%dT%H:%M:%SZ")
|
||||||
|
self.expire_time = expire_time.timestamp() - 60
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"无效的过期时间格式: {expire_str}") from e
|
||||||
|
|
||||||
|
else:
|
||||||
|
self.expire_time = None
|
||||||
|
|
||||||
|
if not self.token:
|
||||||
|
raise ValueError("无法获取有效的访问Token")
|
||||||
|
|
||||||
|
def _is_token_expired(self):
|
||||||
|
"""检查Token是否过期"""
|
||||||
|
if not self.expire_time:
|
||||||
|
return False # 长期Token不过期
|
||||||
|
# 新增调试日志
|
||||||
|
# current_time = time.time()
|
||||||
|
# remaining = self.expire_time - current_time
|
||||||
|
# print(f"Token过期检查: 当前时间 {datetime.fromtimestamp(current_time)} | "
|
||||||
|
# f"过期时间 {datetime.fromtimestamp(self.expire_time)} | "
|
||||||
|
# f"剩余 {remaining:.2f}秒")
|
||||||
|
return time.time() > self.expire_time
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
self.output_file,
|
self.output_file,
|
||||||
@@ -128,6 +162,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
|
if self._is_token_expired():
|
||||||
|
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
|
||||||
|
self._refresh_token()
|
||||||
request_json = {
|
request_json = {
|
||||||
"appkey": self.appkey,
|
"appkey": self.appkey,
|
||||||
"token": self.token,
|
"token": self.token,
|
||||||
@@ -141,12 +178,17 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
|
|
||||||
print(self.api_url, json.dumps(request_json, ensure_ascii=False))
|
print(self.api_url, json.dumps(request_json, ensure_ascii=False))
|
||||||
tmp_file = self.generate_filename()
|
|
||||||
try:
|
try:
|
||||||
resp = requests.post(
|
resp = requests.post(
|
||||||
self.api_url, json.dumps(request_json), headers=self.header
|
self.api_url, json.dumps(request_json), headers=self.header
|
||||||
)
|
)
|
||||||
|
if resp.status_code == 401: # Token过期特殊处理
|
||||||
|
self._refresh_token()
|
||||||
|
resp = requests.post(
|
||||||
|
self.api_url, json.dumps(request_json), headers=self.header
|
||||||
|
)
|
||||||
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
|
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
|
||||||
|
tmp_file = self.generate_filename()
|
||||||
if resp.headers["Content-Type"].startswith("audio/"):
|
if resp.headers["Content-Type"].startswith("audio/"):
|
||||||
with open(tmp_file, "wb") as f:
|
with open(tmp_file, "wb") as f:
|
||||||
f.write(resp.content)
|
f.write(resp.content)
|
||||||
@@ -154,7 +196,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
raise Exception(
|
raise Exception(
|
||||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||||
)
|
)
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
@@ -171,5 +213,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
except FileNotFoundError:
|
except FileNotFoundError:
|
||||||
# 若文件不存在,忽略该异常
|
# 若文件不存在,忽略该异常
|
||||||
pass
|
pass
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise Exception(f"{__name__} error: {e}")
|
raise Exception(f"{__name__} error: {e}")
|
||||||
|
|||||||
@@ -10,12 +10,11 @@ import torch
|
|||||||
import torchaudio
|
import torchaudio
|
||||||
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import os
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import opuslib_next
|
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
|
from core.utils.util import audio_to_data
|
||||||
from core.opus import opus_encoder_utils
|
from core.opus import opus_encoder_utils
|
||||||
import queue
|
import queue
|
||||||
|
|
||||||
@@ -34,7 +33,9 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_audio_queue = queue.Queue()
|
self.tts_audio_queue = queue.Queue()
|
||||||
self.enable_two_way = False
|
self.enable_two_way = False
|
||||||
self.stop_event = threading.Event()
|
self.stop_event = threading.Event()
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(sample_rate=16000, channels=1, frame_size_ms=60)
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
|
)
|
||||||
|
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.punctuations = (
|
self.punctuations = (
|
||||||
@@ -79,12 +80,12 @@ class TTSProviderBase(ABC):
|
|||||||
def _get_segment_text(self):
|
def _get_segment_text(self):
|
||||||
# 合并当前全部文本并处理未分割部分
|
# 合并当前全部文本并处理未分割部分
|
||||||
full_text = "".join(self.tts_text_buff)
|
full_text = "".join(self.tts_text_buff)
|
||||||
current_text = full_text[self.processed_chars:] # 从未处理的位置开始
|
current_text = full_text[self.processed_chars :] # 从未处理的位置开始
|
||||||
last_punct_pos = -1
|
last_punct_pos = -1
|
||||||
for punct in self.punctuations:
|
for punct in self.punctuations:
|
||||||
pos = current_text.rfind(punct)
|
pos = current_text.rfind(punct)
|
||||||
if (pos != -1 and last_punct_pos == -1) or (
|
if (pos != -1 and last_punct_pos == -1) or (
|
||||||
pos != -1 and pos < last_punct_pos
|
pos != -1 and pos < last_punct_pos
|
||||||
):
|
):
|
||||||
last_punct_pos = pos
|
last_punct_pos = pos
|
||||||
if last_punct_pos != -1:
|
if last_punct_pos != -1:
|
||||||
@@ -220,6 +221,7 @@ class TTSProviderBase(ABC):
|
|||||||
)
|
)
|
||||||
self.active_tasks.add(future)
|
self.active_tasks.add(future)
|
||||||
if self.active_tasks:
|
if self.active_tasks:
|
||||||
|
|
||||||
async def wrap_future(future):
|
async def wrap_future(future):
|
||||||
return await asyncio.wrap_future(future)
|
return await asyncio.wrap_future(future)
|
||||||
|
|
||||||
@@ -292,48 +294,13 @@ class TTSProviderBase(ABC):
|
|||||||
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
|
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
|
||||||
raise Exception("该TTS还没有实现stream模式")
|
raise Exception("该TTS还没有实现stream模式")
|
||||||
|
|
||||||
|
def audio_to_pcm_data(self, audio_file_path):
|
||||||
|
"""音频文件转换为PCM编码"""
|
||||||
|
return audio_to_data(audio_file_path, is_opus=False)
|
||||||
|
|
||||||
def audio_to_opus_data(self, audio_file_path):
|
def audio_to_opus_data(self, audio_file_path):
|
||||||
"""音频文件转换为Opus编码"""
|
"""音频文件转换为Opus编码"""
|
||||||
# 获取文件后缀名
|
return audio_to_data(audio_file_path, is_opus=True)
|
||||||
file_type = os.path.splitext(audio_file_path)[1]
|
|
||||||
if file_type:
|
|
||||||
file_type = file_type.lstrip(".")
|
|
||||||
audio = AudioSegment.from_file(audio_file_path, format=file_type)
|
|
||||||
|
|
||||||
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
|
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
|
||||||
|
|
||||||
# 音频时长(秒)
|
|
||||||
duration = len(audio) / 1000.0
|
|
||||||
|
|
||||||
# 获取原始PCM数据(16位小端)
|
|
||||||
raw_data = audio.raw_data
|
|
||||||
|
|
||||||
# 初始化Opus编码器
|
|
||||||
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
|
||||||
|
|
||||||
# 编码参数
|
|
||||||
frame_duration = 60 # 60ms per frame
|
|
||||||
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
|
|
||||||
|
|
||||||
opus_datas = []
|
|
||||||
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
|
||||||
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
|
|
||||||
# 获取当前帧的二进制数据
|
|
||||||
chunk = raw_data[i: i + frame_size * 2]
|
|
||||||
|
|
||||||
# 如果最后一帧不足,补零
|
|
||||||
if len(chunk) < frame_size * 2:
|
|
||||||
chunk += b"\x00" * (frame_size * 2 - len(chunk))
|
|
||||||
|
|
||||||
# 转换为numpy数组处理
|
|
||||||
np_frame = np.frombuffer(chunk, dtype=np.int16)
|
|
||||||
|
|
||||||
# 编码Opus数据
|
|
||||||
opus_data = encoder.encode(np_frame.tobytes(), frame_size)
|
|
||||||
opus_datas.append(opus_data)
|
|
||||||
|
|
||||||
return opus_datas, duration
|
|
||||||
|
|
||||||
def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000):
|
def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000):
|
||||||
tts_speech = torch.from_numpy(
|
tts_speech = torch.from_numpy(
|
||||||
|
|||||||
@@ -16,14 +16,21 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.model = config.get("model")
|
self.model = config.get("model")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.voice = config.get("voice")
|
if config.get("private_voice"):
|
||||||
|
self.voice = config.get("private_voice")
|
||||||
|
else:
|
||||||
|
self.voice = config.get("voice")
|
||||||
|
|
||||||
self.response_format = config.get("response_format")
|
self.response_format = config.get("response_format")
|
||||||
|
|
||||||
self.host = "api.coze.cn"
|
self.host = "api.coze.cn"
|
||||||
self.api_url = f"https://{self.host}/v1/audio/speech"
|
self.api_url = f"https://{self.host}/v1/audio/speech"
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
request_json = {
|
request_json = {
|
||||||
@@ -34,9 +41,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
headers = {
|
headers = {
|
||||||
"Authorization": f"Bearer {self.access_token}",
|
"Authorization": f"Bearer {self.access_token}",
|
||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json",
|
||||||
}
|
}
|
||||||
response = requests.request("POST", self.api_url, json=request_json, headers=headers)
|
response = requests.request(
|
||||||
|
"POST", self.api_url, json=request_json, headers=headers
|
||||||
|
)
|
||||||
data = response.content
|
data = response.content
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
file_to_save = open(tmp_file, "wb")
|
file_to_save = open(tmp_file, "wb")
|
||||||
@@ -45,8 +54,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -22,7 +22,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
|
|
||||||
def generate_filename(self):
|
def generate_filename(self):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
request_params = {}
|
request_params = {}
|
||||||
@@ -37,13 +40,20 @@ class TTSProvider(TTSProviderBase):
|
|||||||
with open(tmp_file, "wb") as file:
|
with open(tmp_file, "wb") as file:
|
||||||
file.write(resp.content)
|
file.write(resp.content)
|
||||||
else:
|
else:
|
||||||
logger.bind(tag=TAG).error(f"Custom TTS请求失败: {resp.status_code} - {resp.text}")
|
logger.bind(tag=TAG).error(
|
||||||
|
f"Custom TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
|
)
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format=self.format)
|
audio = AudioSegment.from_file(tmp_file, format=self.format)
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -17,7 +17,10 @@ logger = setup_logging()
|
|||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.voice = config.get("voice")
|
if config.get("private_voice"):
|
||||||
|
self.voice = config.get("private_voice")
|
||||||
|
else:
|
||||||
|
self.voice = config.get("voice")
|
||||||
|
|
||||||
def generate_filename(self, extension=".mp3"):
|
def generate_filename(self, extension=".mp3"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ class ServeReferenceAudio(BaseModel):
|
|||||||
def decode_audio(cls, values):
|
def decode_audio(cls, values):
|
||||||
audio = values.get("audio")
|
audio = values.get("audio")
|
||||||
if (
|
if (
|
||||||
isinstance(audio, str) and len(audio) > 255
|
isinstance(audio, str) and len(audio) > 255
|
||||||
): # Check if audio is a string (Base64)
|
): # Check if audio is a string (Base64)
|
||||||
try:
|
try:
|
||||||
values["audio"] = base64.b64decode(audio)
|
values["audio"] = base64.b64decode(audio)
|
||||||
@@ -96,32 +96,64 @@ class TTSProvider(TTSProviderBase):
|
|||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
|
|
||||||
self.reference_id = config.get("reference_id")
|
self.reference_id = (
|
||||||
self.reference_audio = config.get("reference_audio", [])
|
None if not config.get("reference_id") else config.get("reference_id")
|
||||||
self.reference_text = config.get("reference_text", [])
|
)
|
||||||
self.format = config.get("format", "wav")
|
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
|
||||||
self.channels = config.get("channels", 1)
|
self.reference_text = parse_string_to_list(config.get("reference_text"))
|
||||||
self.rate = config.get("rate", 44100)
|
self.format = config.get("response_format", "wav")
|
||||||
|
|
||||||
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
||||||
have_key = check_model_key("FishSpeech TTS", self.api_key)
|
have_key = check_model_key("FishSpeech TTS", self.api_key)
|
||||||
if not have_key:
|
if not have_key:
|
||||||
return
|
return
|
||||||
self.normalize = config.get("normalize", True)
|
self.normalize = str(config.get("normalize", True)).lower() in (
|
||||||
self.max_new_tokens = config.get("max_new_tokens", 1024)
|
"true",
|
||||||
self.chunk_length = config.get("chunk_length", 200)
|
"1",
|
||||||
self.top_p = config.get("top_p", 0.7)
|
"yes",
|
||||||
self.repetition_penalty = config.get("repetition_penalty", 1.2)
|
)
|
||||||
self.temperature = config.get("temperature", 0.7)
|
|
||||||
self.streaming = config.get("streaming", False)
|
# 处理空字符串的情况
|
||||||
|
channels = config.get("channels", "1")
|
||||||
|
rate = config.get("rate", "44100")
|
||||||
|
max_new_tokens = config.get("max_new_tokens", "1024")
|
||||||
|
chunk_length = config.get("chunk_length", "200")
|
||||||
|
|
||||||
|
self.channels = int(channels) if channels else 1
|
||||||
|
self.rate = int(rate) if rate else 44100
|
||||||
|
self.max_new_tokens = int(max_new_tokens) if max_new_tokens else 1024
|
||||||
|
self.chunk_length = int(chunk_length) if chunk_length else 200
|
||||||
|
|
||||||
|
# 处理空字符串的情况
|
||||||
|
top_p = config.get("top_p", "0.7")
|
||||||
|
temperature = config.get("temperature", "0.7")
|
||||||
|
repetition_penalty = config.get("repetition_penalty", "1.2")
|
||||||
|
|
||||||
|
self.top_p = float(top_p) if top_p else 0.7
|
||||||
|
self.temperature = float(temperature) if temperature else 0.7
|
||||||
|
self.repetition_penalty = (
|
||||||
|
float(repetition_penalty) if repetition_penalty else 1.2
|
||||||
|
)
|
||||||
|
|
||||||
|
self.streaming = str(config.get("streaming", False)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
self.use_memory_cache = config.get("use_memory_cache", "on")
|
self.use_memory_cache = config.get("use_memory_cache", "on")
|
||||||
self.seed = config.get("seed")
|
self.seed = int(config.get("seed")) if config.get("seed") else None
|
||||||
self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts")
|
self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts")
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
def _get_audio_from_tts(self, data_bytes):
|
def _get_audio_from_tts(self, data_bytes):
|
||||||
tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0)
|
tts_speech = torch.from_numpy(
|
||||||
|
np.array(np.frombuffer(data_bytes, dtype=np.int16))
|
||||||
|
).unsqueeze(dim=0)
|
||||||
with io.BytesIO() as bf:
|
with io.BytesIO() as bf:
|
||||||
torchaudio.save(bf, tts_speech, 44100, format="wav")
|
torchaudio.save(bf, tts_speech, 44100, format="wav")
|
||||||
audio = AudioSegment.from_file(bf, format="wav")
|
audio = AudioSegment.from_file(bf, format="wav")
|
||||||
@@ -147,29 +179,35 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
# Prepare reference data
|
# Prepare reference data
|
||||||
if self.reference_audio and self.reference_text:
|
if self.reference_audio and self.reference_text:
|
||||||
byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio]
|
byte_audios = [
|
||||||
ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text]
|
audio_to_bytes(ref_audio) for ref_audio in self.reference_audio
|
||||||
data["references"] = [
|
]
|
||||||
ServeReferenceAudio(
|
ref_texts = [
|
||||||
audio=audio if audio else b"", text=text
|
read_ref_text(ref_text) for ref_text in self.reference_text
|
||||||
)
|
]
|
||||||
for text, audio in zip(ref_texts, byte_audios)
|
data["references"] = (
|
||||||
],
|
[
|
||||||
|
ServeReferenceAudio(audio=audio if audio else b"", text=text)
|
||||||
|
for text, audio in zip(ref_texts, byte_audios)
|
||||||
|
],
|
||||||
|
)
|
||||||
data["reference_id"] = None
|
data["reference_id"] = None
|
||||||
|
|
||||||
pydantic_data = ServeTTSRequest(**data)
|
pydantic_data = ServeTTSRequest(**data)
|
||||||
audio_buff = None
|
audio_buff = None
|
||||||
chunk_total = b''
|
chunk_total = b""
|
||||||
last_raw = b''
|
last_raw = b""
|
||||||
audio_raw = b''
|
audio_raw = b""
|
||||||
print("请求tts")
|
print("请求tts")
|
||||||
with requests.post(
|
with requests.post(
|
||||||
self.api_url,
|
self.api_url,
|
||||||
data=ormsgpack.packb(pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC),
|
data=ormsgpack.packb(
|
||||||
headers={
|
pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC
|
||||||
"Authorization": f"Bearer {self.api_key}",
|
),
|
||||||
"Content-Type": "application/msgpack",
|
headers={
|
||||||
},
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
|
"Content-Type": "application/msgpack",
|
||||||
|
},
|
||||||
) as response:
|
) as response:
|
||||||
if response.status_code == 200:
|
if response.status_code == 200:
|
||||||
index = 0
|
index = 0
|
||||||
@@ -177,34 +215,55 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 拼接当前块和上一块数据
|
# 拼接当前块和上一块数据
|
||||||
chunk_total += chunk
|
chunk_total += chunk
|
||||||
# 最后一个是静音,说明是一个完整的音频
|
# 最后一个是静音,说明是一个完整的音频
|
||||||
if len(chunk_total) % 2 == 0 and chunk_total[-2:] == b'\x00\x00':
|
if (
|
||||||
|
len(chunk_total) % 2 == 0
|
||||||
|
and chunk_total[-2:] == b"\x00\x00"
|
||||||
|
):
|
||||||
audio = self._get_audio_from_tts(chunk_total)
|
audio = self._get_audio_from_tts(chunk_total)
|
||||||
audio_raw = audio_raw + audio.raw_data
|
audio_raw = audio_raw + audio.raw_data
|
||||||
# 长度凑够2贞开始发送,60ms*2=120ms
|
# 长度凑够2贞开始发送,60ms*2=120ms
|
||||||
if len(audio_raw) >= 3840:
|
if len(audio_raw) >= 3840:
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
||||||
if index == 0:
|
if index == 0:
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
|
yield TTSMessageDTO(
|
||||||
content=opus_datas,
|
u_id=u_id,
|
||||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_START)
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
|
yield TTSMessageDTO(
|
||||||
content=opus_datas,
|
u_id=u_id,
|
||||||
tts_finish_text=text, sentence_type=None)
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
audio_raw = b''
|
content=opus_datas,
|
||||||
chunk_total = b''
|
tts_finish_text=text,
|
||||||
|
sentence_type=None,
|
||||||
|
)
|
||||||
|
audio_raw = b""
|
||||||
|
chunk_total = b""
|
||||||
if len(chunk_total) > 0:
|
if len(chunk_total) > 0:
|
||||||
audio = self._get_audio_from_tts(chunk_total)
|
audio = self._get_audio_from_tts(chunk_total)
|
||||||
audio_raw = audio_raw + audio.raw_data
|
audio_raw = audio_raw + audio.raw_data
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas,
|
yield TTSMessageDTO(
|
||||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_END,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[],
|
yield TTSMessageDTO(
|
||||||
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=[],
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_END,
|
||||||
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
print('请求失败:', response.status_code, response.text)
|
print("请求失败:", response.status_code, response.text)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error("tts发生错误")
|
logger.bind(tag=TAG).error("tts发生错误")
|
||||||
traceback.print_exc()
|
traceback.print_exc()
|
||||||
|
|||||||
@@ -1,10 +1,8 @@
|
|||||||
import os
|
import os
|
||||||
import uuid
|
import uuid
|
||||||
import json
|
|
||||||
import base64
|
|
||||||
import requests
|
import requests
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -13,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
@@ -21,23 +20,62 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.ref_audio_path = config.get("ref_audio_path")
|
self.ref_audio_path = config.get("ref_audio_path")
|
||||||
self.prompt_text = config.get("prompt_text")
|
self.prompt_text = config.get("prompt_text")
|
||||||
self.prompt_lang = config.get("prompt_lang", "zh")
|
self.prompt_lang = config.get("prompt_lang", "zh")
|
||||||
self.top_k = config.get("top_k", 5)
|
|
||||||
self.top_p = config.get("top_p", 1)
|
# 处理空字符串的情况
|
||||||
self.temperature = config.get("temperature", 1)
|
top_k = config.get("top_k", "5")
|
||||||
|
top_p = config.get("top_p", "1")
|
||||||
|
temperature = config.get("temperature", "1")
|
||||||
|
batch_threshold = config.get("batch_threshold", "0.75")
|
||||||
|
batch_size = config.get("batch_size", "1")
|
||||||
|
speed_factor = config.get("speed_factor", "1.0")
|
||||||
|
seed = config.get("seed", "-1")
|
||||||
|
repetition_penalty = config.get("repetition_penalty", "1.35")
|
||||||
|
|
||||||
|
self.top_k = int(top_k) if top_k else 5
|
||||||
|
self.top_p = float(top_p) if top_p else 1
|
||||||
|
self.temperature = float(temperature) if temperature else 1
|
||||||
|
self.batch_threshold = float(batch_threshold) if batch_threshold else 0.75
|
||||||
|
self.batch_size = int(batch_size) if batch_size else 1
|
||||||
|
self.speed_factor = float(speed_factor) if speed_factor else 1.0
|
||||||
|
self.seed = int(seed) if seed else -1
|
||||||
|
self.repetition_penalty = (
|
||||||
|
float(repetition_penalty) if repetition_penalty else 1.35
|
||||||
|
)
|
||||||
|
|
||||||
self.text_split_method = config.get("text_split_method", "cut0")
|
self.text_split_method = config.get("text_split_method", "cut0")
|
||||||
self.batch_size = config.get("batch_size", 1)
|
|
||||||
self.batch_threshold = config.get("batch_threshold", 0.75)
|
self.split_bucket = str(config.get("split_bucket", True)).lower() in (
|
||||||
self.split_bucket = config.get("split_bucket", True)
|
"true",
|
||||||
self.return_fragment = config.get("return_fragment", False)
|
"1",
|
||||||
self.speed_factor = config.get("speed_factor", 1.0)
|
"yes",
|
||||||
self.streaming_mode = config.get("streaming_mode", False)
|
)
|
||||||
self.seed = config.get("seed", -1)
|
self.return_fragment = str(config.get("return_fragment", False)).lower() in (
|
||||||
self.parallel_infer = config.get("parallel_infer", True)
|
"true",
|
||||||
self.repetition_penalty = config.get("repetition_penalty", 1.35)
|
"1",
|
||||||
self.aux_ref_audio_paths = config.get("aux_ref_audio_paths", [])
|
"yes",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.streaming_mode = str(config.get("streaming_mode", False)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.parallel_infer = str(config.get("parallel_infer", True)).lower() in (
|
||||||
|
"true",
|
||||||
|
"1",
|
||||||
|
"yes",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.aux_ref_audio_paths = parse_string_to_list(
|
||||||
|
config.get("aux_ref_audio_paths")
|
||||||
|
)
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
@@ -60,7 +98,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"streaming_mode": self.streaming_mode,
|
"streaming_mode": self.streaming_mode,
|
||||||
"seed": self.seed,
|
"seed": self.seed,
|
||||||
"parallel_infer": self.parallel_infer,
|
"parallel_infer": self.parallel_infer,
|
||||||
"repetition_penalty": self.repetition_penalty
|
"repetition_penalty": self.repetition_penalty,
|
||||||
}
|
}
|
||||||
|
|
||||||
resp = requests.post(self.url, json=request_json)
|
resp = requests.post(self.url, json=request_json)
|
||||||
@@ -68,13 +106,20 @@ class TTSProvider(TTSProviderBase):
|
|||||||
with open(tmp_file, "wb") as file:
|
with open(tmp_file, "wb") as file:
|
||||||
file.write(resp.content)
|
file.write(resp.content)
|
||||||
else:
|
else:
|
||||||
logger.bind(tag=TAG).error(f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}")
|
logger.bind(tag=TAG).error(
|
||||||
|
f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
|
)
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import os
|
|||||||
import uuid
|
import uuid
|
||||||
import requests
|
import requests
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -11,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
@@ -19,18 +20,29 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.prompt_text = config.get("prompt_text")
|
self.prompt_text = config.get("prompt_text")
|
||||||
self.prompt_language = config.get("prompt_language")
|
self.prompt_language = config.get("prompt_language")
|
||||||
self.text_language = config.get("text_language", "audo")
|
self.text_language = config.get("text_language", "audo")
|
||||||
self.top_k = config.get("top_k", 15)
|
|
||||||
self.top_p = config.get("top_p", 1.0)
|
|
||||||
self.temperature = config.get("temperature", 1.0)
|
|
||||||
self.cut_punc = config.get("cut_punc","")
|
|
||||||
self.speed = config.get("speed", 1.0)
|
|
||||||
self.inp_refs = config.get("inp_refs",[])
|
|
||||||
self.sample_steps = config.get("sample_steps",32)
|
|
||||||
self.if_sr = config.get("if_sr",False)
|
|
||||||
|
|
||||||
|
# 处理空字符串的情况
|
||||||
|
top_k = config.get("top_k", "15")
|
||||||
|
top_p = config.get("top_p", "1.0")
|
||||||
|
temperature = config.get("temperature", "1.0")
|
||||||
|
sample_steps = config.get("sample_steps", "32")
|
||||||
|
speed = config.get("speed", "1.0")
|
||||||
|
|
||||||
|
self.top_k = int(top_k) if top_k else 15
|
||||||
|
self.top_p = float(top_p) if top_p else 1.0
|
||||||
|
self.temperature = float(temperature) if temperature else 1.0
|
||||||
|
self.sample_steps = int(sample_steps) if sample_steps else 32
|
||||||
|
self.speed = float(speed) if speed else 1.0
|
||||||
|
|
||||||
|
self.cut_punc = config.get("cut_punc", "")
|
||||||
|
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
|
||||||
|
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
@@ -55,13 +67,20 @@ class TTSProvider(TTSProviderBase):
|
|||||||
with open(tmp_file, "wb") as file:
|
with open(tmp_file, "wb") as file:
|
||||||
file.write(resp.content)
|
file.write(resp.content)
|
||||||
else:
|
else:
|
||||||
logger.bind(tag=TAG).error(f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}")
|
logger.bind(tag=TAG).error(
|
||||||
|
f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
|
)
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -2,10 +2,9 @@ import os
|
|||||||
import uuid
|
import uuid
|
||||||
import json
|
import json
|
||||||
import requests
|
import requests
|
||||||
|
from core.utils.util import parse_string_to_list
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
|
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
||||||
|
|
||||||
@@ -16,30 +15,35 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.group_id = config.get("group_id")
|
self.group_id = config.get("group_id")
|
||||||
self.api_key = config.get("api_key")
|
self.api_key = config.get("api_key")
|
||||||
self.model = config.get("model")
|
self.model = config.get("model")
|
||||||
self.voice_id = config.get("voice_id")
|
if config.get("private_voice"):
|
||||||
|
self.voice_id = config.get("private_voice")
|
||||||
|
else:
|
||||||
|
self.voice_id = config.get("voice_id")
|
||||||
|
|
||||||
default_voice_setting = {
|
default_voice_setting = {
|
||||||
"voice_id": "female-shaonv",
|
"voice_id": "female-shaonv",
|
||||||
"speed": 1,
|
"speed": 1,
|
||||||
"vol": 1,
|
"vol": 1,
|
||||||
"pitch": 0,
|
"pitch": 0,
|
||||||
"emotion": "happy"
|
"emotion": "happy",
|
||||||
}
|
|
||||||
default_pronunciation_dict = {
|
|
||||||
"tone": [
|
|
||||||
"处理/(chu3)(li3)", "危险/dangerous"
|
|
||||||
]
|
|
||||||
}
|
}
|
||||||
|
default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]}
|
||||||
defult_audio_setting = {
|
defult_audio_setting = {
|
||||||
"sample_rate": 32000,
|
"sample_rate": 32000,
|
||||||
"bitrate": 128000,
|
"bitrate": 128000,
|
||||||
"format": "mp3",
|
"format": "mp3",
|
||||||
"channel": 1
|
"channel": 1,
|
||||||
|
}
|
||||||
|
self.voice_setting = {
|
||||||
|
**default_voice_setting,
|
||||||
|
**config.get("voice_setting", {}),
|
||||||
|
}
|
||||||
|
self.pronunciation_dict = {
|
||||||
|
**default_pronunciation_dict,
|
||||||
|
**config.get("pronunciation_dict", {}),
|
||||||
}
|
}
|
||||||
self.voice_setting = {**default_voice_setting, **config.get("voice_setting", {})}
|
|
||||||
self.pronunciation_dict = {**default_pronunciation_dict, **config.get("pronunciation_dict", {})}
|
|
||||||
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
||||||
self.timber_weights = config.get("timber_weights", [])
|
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
|
||||||
|
|
||||||
if self.voice_id:
|
if self.voice_id:
|
||||||
self.voice_setting["voice_id"] = self.voice_id
|
self.voice_setting["voice_id"] = self.voice_id
|
||||||
@@ -48,11 +52,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
|
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
|
||||||
self.header = {
|
self.header = {
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
"Authorization": f"Bearer {self.api_key}"
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
}
|
}
|
||||||
|
|
||||||
def generate_filename(self, extension=".mp3"):
|
def generate_filename(self, extension=".mp3"):
|
||||||
return os.path.join(self.output_file, f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
@@ -70,25 +77,34 @@ class TTSProvider(TTSProviderBase):
|
|||||||
request_json["voice_setting"]["voice_id"] = ""
|
request_json["voice_setting"]["voice_id"] = ""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header)
|
resp = requests.post(
|
||||||
|
self.api_url, json.dumps(request_json), headers=self.header
|
||||||
|
)
|
||||||
# 检查返回请求数据的status_code是否为0
|
# 检查返回请求数据的status_code是否为0
|
||||||
if resp.json()["base_resp"]["status_code"] == 0:
|
if resp.json()["base_resp"]["status_code"] == 0:
|
||||||
data = resp.json()['data']['audio']
|
data = resp.json()["data"]["audio"]
|
||||||
file_to_save = open(tmp_file, "wb")
|
file_to_save = open(tmp_file, "wb")
|
||||||
file_to_save.write(bytes.fromhex(data))
|
file_to_save.write(bytes.fromhex(data))
|
||||||
else:
|
else:
|
||||||
raise Exception(f"{__name__} status_code: {resp.status_code} response: {resp.content}")
|
raise Exception(
|
||||||
|
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise Exception(f"{__name__} error: {e}")
|
raise Exception(f"{__name__} error: {e}")
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format="mp3")
|
audio = AudioSegment.from_file(tmp_file, format="mp3")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
except FileNotFoundError:
|
except FileNotFoundError:
|
||||||
# 若文件不存在,忽略该异常
|
# 若文件不存在,忽略该异常
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -9,46 +9,64 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
|||||||
from core.utils.util import check_model_key
|
from core.utils.util import check_model_key
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
|
||||||
|
|
||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.api_key = config.get("api_key")
|
self.api_key = config.get("api_key")
|
||||||
self.api_url = config.get("api_url", "https://api.openai.com/v1/audio/speech")
|
self.api_url = config.get("api_url", "https://api.openai.com/v1/audio/speech")
|
||||||
self.model = config.get("model", "tts-1")
|
self.model = config.get("model", "tts-1")
|
||||||
self.voice = config.get("voice", "alloy")
|
if config.get("private_voice"):
|
||||||
|
self.voice = config.get("private_voice")
|
||||||
|
else:
|
||||||
|
self.voice = config.get("voice", "alloy")
|
||||||
self.response_format = "wav"
|
self.response_format = "wav"
|
||||||
self.speed = config.get("speed", 1.0)
|
|
||||||
|
# 处理空字符串的情况
|
||||||
|
speed = config.get("speed", "1.0")
|
||||||
|
self.speed = float(speed) if speed else 1.0
|
||||||
|
|
||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
check_model_key("TTS", self.api_key)
|
check_model_key("TTS", self.api_key)
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
headers = {
|
headers = {
|
||||||
"Authorization": f"Bearer {self.api_key}",
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json",
|
||||||
}
|
}
|
||||||
data = {
|
data = {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
"input": text,
|
"input": text,
|
||||||
"voice": self.voice,
|
"voice": self.voice,
|
||||||
"response_format": "wav",
|
"response_format": "wav",
|
||||||
"speed": self.speed
|
"speed": self.speed,
|
||||||
}
|
}
|
||||||
response = requests.post(self.api_url, json=data, headers=headers)
|
response = requests.post(self.api_url, json=data, headers=headers)
|
||||||
if response.status_code == 200:
|
if response.status_code == 200:
|
||||||
with open(tmp_file, "wb") as audio_file:
|
with open(tmp_file, "wb") as audio_file:
|
||||||
audio_file.write(response.content)
|
audio_file.write(response.content)
|
||||||
else:
|
else:
|
||||||
raise Exception(f"OpenAI TTS请求失败: {response.status_code} - {response.text}")
|
raise Exception(
|
||||||
|
f"OpenAI TTS请求失败: {response.status_code} - {response.text}"
|
||||||
|
)
|
||||||
# 使用 pydub 读取临时文件
|
# 使用 pydub 读取临时文件
|
||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -14,17 +14,23 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.model = config.get("model")
|
self.model = config.get("model")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.voice = config.get("voice")
|
if config.get("private_voice"):
|
||||||
|
self.voice = config.get("private_voice")
|
||||||
|
else:
|
||||||
|
self.voice = config.get("voice")
|
||||||
self.response_format = config.get("response_format")
|
self.response_format = config.get("response_format")
|
||||||
self.sample_rate = config.get("sample_rate")
|
self.sample_rate = config.get("sample_rate")
|
||||||
self.speed = config.get("speed")
|
self.speed = float(config.get("speed", 1.0))
|
||||||
self.gain = config.get("gain")
|
self.gain = config.get("gain")
|
||||||
|
|
||||||
self.host = "api.siliconflow.cn"
|
self.host = "api.siliconflow.cn"
|
||||||
self.api_url = f"https://{self.host}/v1/audio/speech"
|
self.api_url = f"https://{self.host}/v1/audio/speech"
|
||||||
|
|
||||||
def generate_filename(self, extension=".wav"):
|
def generate_filename(self, extension=".wav"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
@@ -36,9 +42,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
headers = {
|
headers = {
|
||||||
"Authorization": f"Bearer {self.access_token}",
|
"Authorization": f"Bearer {self.access_token}",
|
||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json",
|
||||||
}
|
}
|
||||||
response = requests.request("POST", self.api_url, json=request_json, headers=headers)
|
response = requests.request(
|
||||||
|
"POST", self.api_url, json=request_json, headers=headers
|
||||||
|
)
|
||||||
data = response.content
|
data = response.content
|
||||||
file_to_save = open(tmp_file, "wb")
|
file_to_save = open(tmp_file, "wb")
|
||||||
file_to_save.write(data)
|
file_to_save.write(data)
|
||||||
@@ -46,8 +54,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
audio = AudioSegment.from_file(tmp_file, format="wav")
|
audio = AudioSegment.from_file(tmp_file, format="wav")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
|
|||||||
@@ -14,49 +14,63 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
|
|||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.url = config.get("url", "https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=")
|
self.url = config.get(
|
||||||
self.voice_id = config.get("voice_id", 1695)
|
"url",
|
||||||
|
"https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=",
|
||||||
|
)
|
||||||
|
if config.get("private_voice"):
|
||||||
|
self.voice_id = int(config.get("private_voice"))
|
||||||
|
else:
|
||||||
|
self.voice_id = int(config.get("voice_id", 1695))
|
||||||
self.token = config.get("token")
|
self.token = config.get("token")
|
||||||
self.to_lang = config.get("to_lang")
|
self.to_lang = config.get("to_lang")
|
||||||
self.volume_change_dB = config.get("volume_change_dB", 0)
|
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
||||||
self.speed_factor = config.get("speed_factor", 1)
|
self.speed_factor = int(config.get("speed_factor", 1))
|
||||||
self.stream = config.get("stream", False)
|
self.stream = str(config.get("stream", False)).lower() in ("true", "1", "yes")
|
||||||
self.output_file = config.get("output_dir")
|
self.output_file = config.get("output_dir")
|
||||||
self.pitch_factor = config.get("pitch_factor", 0)
|
self.pitch_factor = int(config.get("pitch_factor", 0))
|
||||||
self.format = config.get("format", "mp3")
|
self.format = config.get("format", "mp3")
|
||||||
self.emotion = config.get("emotion", 1)
|
self.emotion = int(config.get("emotion", 1))
|
||||||
self.header = {
|
self.header = {"Content-Type": "application/json"}
|
||||||
"Content-Type": "application/json"
|
|
||||||
}
|
|
||||||
|
|
||||||
def generate_filename(self, extension=".mp3"):
|
def generate_filename(self, extension=".mp3"):
|
||||||
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
|
return os.path.join(
|
||||||
|
self.output_file,
|
||||||
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
|
)
|
||||||
|
|
||||||
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False):
|
||||||
tmp_file = self.generate_filename()
|
tmp_file = self.generate_filename()
|
||||||
url = f'{self.url}{self.token}'
|
url = f"{self.url}{self.token}"
|
||||||
result = "firefly"
|
result = "firefly"
|
||||||
payload = json.dumps({
|
payload = json.dumps(
|
||||||
"to_lang": self.to_lang,
|
{
|
||||||
"text": text,
|
"to_lang": self.to_lang,
|
||||||
"emotion": self.emotion,
|
"text": text,
|
||||||
"format": self.format,
|
"emotion": self.emotion,
|
||||||
"volume_change_dB": self.volume_change_dB,
|
"format": self.format,
|
||||||
"voice_id": self.voice_id,
|
"volume_change_dB": self.volume_change_dB,
|
||||||
"pitch_factor": self.pitch_factor,
|
"voice_id": self.voice_id,
|
||||||
"speed_factor": self.speed_factor,
|
"pitch_factor": self.pitch_factor,
|
||||||
"token": self.token
|
"speed_factor": self.speed_factor,
|
||||||
})
|
"token": self.token,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
resp = requests.request("POST", url, data=payload)
|
resp = requests.request("POST", url, data=payload)
|
||||||
if resp.status_code != 200:
|
if resp.status_code != 200:
|
||||||
return
|
return
|
||||||
resp_json = resp.json()
|
resp_json = resp.json()
|
||||||
try:
|
try:
|
||||||
result = resp_json['url'] + ':' + str(
|
result = (
|
||||||
resp_json[
|
resp_json["url"]
|
||||||
'port']) + '/flashsummary/retrieveFileData?stream=True&token=' + self.token + '&voice_audio_path=' + \
|
+ ":"
|
||||||
resp_json['voice_path']
|
+ str(resp_json["port"])
|
||||||
|
+ "/flashsummary/retrieveFileData?stream=True&token="
|
||||||
|
+ self.token
|
||||||
|
+ "&voice_audio_path="
|
||||||
|
+ resp_json["voice_path"]
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print("error:", e)
|
print("error:", e)
|
||||||
|
|
||||||
@@ -67,8 +81,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
audio = AudioSegment.from_file(tmp_file, format="mp3")
|
audio = AudioSegment.from_file(tmp_file, format="mp3")
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000)
|
audio = audio.set_channels(1).set_frame_rate(16000)
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
|
||||||
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text,
|
yield TTSMessageDTO(
|
||||||
sentence_type=SentenceType.SENTENCE_START)
|
u_id=u_id,
|
||||||
|
msg_type=MsgType.TTS_TEXT_RESPONSE,
|
||||||
|
content=opus_datas,
|
||||||
|
tts_finish_text=text,
|
||||||
|
sentence_type=SentenceType.SENTENCE_START,
|
||||||
|
)
|
||||||
# 用完后删除临时文件
|
# 用完后删除临时文件
|
||||||
try:
|
try:
|
||||||
os.remove(tmp_file)
|
os.remove(tmp_file)
|
||||||
@@ -78,4 +97,3 @@ class TTSProvider(TTSProviderBase):
|
|||||||
voice_path = resp_json.get("voice_path")
|
voice_path = resp_json.get("voice_path")
|
||||||
des_path = tmp_file
|
des_path = tmp_file
|
||||||
shutil.move(voice_path, des_path)
|
shutil.move(voice_path, des_path)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user