update:合并main分支

This commit is contained in:
hrz
2025-05-21 14:52:24 +08:00
parent 191ac47353
commit c4c84e44e1
13 changed files with 458 additions and 233 deletions
@@ -239,7 +239,7 @@ class ASRProvider(ASRProviderBase):
) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本"""
if self._is_token_expired():
logger.warning("Token已过期,正在自动刷新...")
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
self._refresh_token()
file_path = None
@@ -6,18 +6,18 @@ import hashlib
import base64
import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
from pydub import AudioSegment
from core.providers.tts.base import TTSProviderBase
import http.client
import urllib.parse
import time
import uuid
from urllib import parse
from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class AccessToken:
@@ -48,7 +48,7 @@ class AccessToken:
}
# 构造规范化的请求字符串
query_string = AccessToken._encode_dict(parameters)
print("规范化的请求字符串: %s" % query_string)
# print('规范化的请求字符串: %s' % query_string)
# 构造待签名字符串
string_to_sign = (
"GET"
@@ -57,7 +57,7 @@ class AccessToken:
+ "&"
+ AccessToken._encode_text(query_string)
)
print("待签名的字符串: %s" % string_to_sign)
# print('待签名的字符串: %s' % string_to_sign)
# 计算签名
secreted_string = hmac.new(
bytes(access_key_secret + "&", encoding="utf-8"),
@@ -65,16 +65,16 @@ class AccessToken:
hashlib.sha1,
).digest()
signature = base64.b64encode(secreted_string)
print("签名: %s" % signature)
# print('签名: %s' % signature)
# 进行URL编码
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" % (
signature,
query_string,
)
print("url: %s" % full_url)
# print('url: %s' % full_url)
# 提交HTTP GET请求
response = requests.get(full_url)
if response.ok:
@@ -84,7 +84,7 @@ class AccessToken:
token = root_obj[key]["Id"]
expire_time = root_obj[key]["ExpireTime"]
return token, expire_time
print(response.text)
# print(response.text)
return None, None
@@ -94,33 +94,67 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
# 新增空值判断逻辑
access_key_id = config.get("access_key_id")
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.access_key_id = config.get("access_key_id")
self.access_key_secret = config.get("access_key_secret")
self.appkey = config.get("appkey")
self.token = token
self.format = config.get("format", "wav")
self.sample_rate = config.get("sample_rate", 16000)
self.voice = config.get("voice", "xiaoyun")
self.volume = config.get("volume", 50)
self.speech_rate = config.get("speech_rate", 0)
self.pitch_rate = config.get("pitch_rate", 0)
self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com")
self.api_url = f"https://{self.host}/stream/v1/tts"
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"):
return os.path.join(
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):
if self._is_token_expired():
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
self._refresh_token()
request_json = {
"appkey": self.appkey,
"token": self.token,
@@ -141,12 +178,17 @@ class TTSProvider(TTSProviderBase):
}
print(self.api_url, json.dumps(request_json, ensure_ascii=False))
tmp_file = self.generate_filename()
try:
resp = requests.post(
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格式的
tmp_file = self.generate_filename()
if resp.headers["Content-Type"].startswith("audio/"):
with open(tmp_file, "wb") as f:
f.write(resp.content)
@@ -154,7 +196,7 @@ class TTSProvider(TTSProviderBase):
raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
)
# 使用 pydub 读取临时文件
# 使用 pydub 读取临时文件
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data)
@@ -171,5 +213,6 @@ class TTSProvider(TTSProviderBase):
except FileNotFoundError:
# 若文件不存在,忽略该异常
pass
except Exception as e:
raise Exception(f"{__name__} error: {e}")
+12 -45
View File
@@ -10,12 +10,11 @@ import torch
import torchaudio
from config.logger import setup_logging
import os
import numpy as np
import opuslib_next
from pydub import AudioSegment
from abc import ABC, abstractmethod
from core.utils import textUtils
from core.utils.util import audio_to_data
from core.opus import opus_encoder_utils
import queue
@@ -34,7 +33,9 @@ class TTSProviderBase(ABC):
self.tts_audio_queue = queue.Queue()
self.enable_two_way = False
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.punctuations = (
@@ -79,12 +80,12 @@ class TTSProviderBase(ABC):
def _get_segment_text(self):
# 合并当前全部文本并处理未分割部分
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
for punct in self.punctuations:
pos = current_text.rfind(punct)
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
if last_punct_pos != -1:
@@ -220,6 +221,7 @@ class TTSProviderBase(ABC):
)
self.active_tasks.add(future)
if self.active_tasks:
async def 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):
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):
"""音频文件转换为Opus编码"""
# 获取文件后缀名
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
return audio_to_data(audio_file_path, is_opus=True)
def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000):
tts_speech = torch.from_numpy(
@@ -16,14 +16,21 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.model = config.get("model")
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.host = "api.coze.cn"
self.api_url = f"https://{self.host}/v1/audio/speech"
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):
request_json = {
@@ -34,9 +41,11 @@ class TTSProvider(TTSProviderBase):
}
headers = {
"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
tmp_file = self.generate_filename()
file_to_save = open(tmp_file, "wb")
@@ -45,8 +54,13 @@ class TTSProvider(TTSProviderBase):
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -22,7 +22,10 @@ class TTSProvider(TTSProviderBase):
self.output_file = config.get("output_dir", "tmp/")
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):
request_params = {}
@@ -37,13 +40,20 @@ class TTSProvider(TTSProviderBase):
with open(tmp_file, "wb") as file:
file.write(resp.content)
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 读取临时文件
audio = AudioSegment.from_file(tmp_file, format=self.format)
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -17,7 +17,10 @@ logger = setup_logging()
class TTSProvider(TTSProviderBase):
def __init__(self, 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"):
return os.path.join(
@@ -35,7 +35,7 @@ class ServeReferenceAudio(BaseModel):
def decode_audio(cls, values):
audio = values.get("audio")
if (
isinstance(audio, str) and len(audio) > 255
isinstance(audio, str) and len(audio) > 255
): # Check if audio is a string (Base64)
try:
values["audio"] = base64.b64decode(audio)
@@ -96,32 +96,64 @@ class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.reference_id = config.get("reference_id")
self.reference_audio = config.get("reference_audio", [])
self.reference_text = config.get("reference_text", [])
self.format = config.get("format", "wav")
self.channels = config.get("channels", 1)
self.rate = config.get("rate", 44100)
self.reference_id = (
None if not config.get("reference_id") else config.get("reference_id")
)
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
self.reference_text = parse_string_to_list(config.get("reference_text"))
self.format = config.get("response_format", "wav")
self.api_key = config.get("api_key", "YOUR_API_KEY")
have_key = check_model_key("FishSpeech TTS", self.api_key)
if not have_key:
return
self.normalize = config.get("normalize", True)
self.max_new_tokens = config.get("max_new_tokens", 1024)
self.chunk_length = config.get("chunk_length", 200)
self.top_p = config.get("top_p", 0.7)
self.repetition_penalty = config.get("repetition_penalty", 1.2)
self.temperature = config.get("temperature", 0.7)
self.streaming = config.get("streaming", False)
self.normalize = str(config.get("normalize", True)).lower() in (
"true",
"1",
"yes",
)
# 处理空字符串的情况
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.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")
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):
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:
torchaudio.save(bf, tts_speech, 44100, format="wav")
audio = AudioSegment.from_file(bf, format="wav")
@@ -147,29 +179,35 @@ class TTSProvider(TTSProviderBase):
# Prepare reference data
if self.reference_audio and self.reference_text:
byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio]
ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text]
data["references"] = [
ServeReferenceAudio(
audio=audio if audio else b"", text=text
)
for text, audio in zip(ref_texts, byte_audios)
],
byte_audios = [
audio_to_bytes(ref_audio) for ref_audio in self.reference_audio
]
ref_texts = [
read_ref_text(ref_text) for ref_text in self.reference_text
]
data["references"] = (
[
ServeReferenceAudio(audio=audio if audio else b"", text=text)
for text, audio in zip(ref_texts, byte_audios)
],
)
data["reference_id"] = None
pydantic_data = ServeTTSRequest(**data)
audio_buff = None
chunk_total = b''
last_raw = b''
audio_raw = b''
chunk_total = b""
last_raw = b""
audio_raw = b""
print("请求tts")
with requests.post(
self.api_url,
data=ormsgpack.packb(pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC),
headers={
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/msgpack",
},
self.api_url,
data=ormsgpack.packb(
pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC
),
headers={
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/msgpack",
},
) as response:
if response.status_code == 200:
index = 0
@@ -177,34 +215,55 @@ class TTSProvider(TTSProviderBase):
# 拼接当前块和上一块数据
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_raw = audio_raw + audio.raw_data
# 长度凑够2贞开始发送,60ms*2=120ms
if len(audio_raw) >= 3840:
opus_datas = self.wav_to_opus_data_audio_raw(audio_raw)
if index == 0:
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
else:
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text, sentence_type=None)
audio_raw = b''
chunk_total = b''
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=None,
)
audio_raw = b""
chunk_total = b""
if len(chunk_total) > 0:
audio = self._get_audio_from_tts(chunk_total)
audio_raw = audio_raw + audio.raw_data
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,
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_END,
)
else:
yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[],
tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=[],
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_END,
)
else:
print('请求失败:', response.status_code, response.text)
print("请求失败:", response.status_code, response.text)
except Exception as e:
logger.bind(tag=TAG).error("tts发生错误")
traceback.print_exc()
@@ -1,10 +1,8 @@
import os
import uuid
import json
import base64
import requests
from pydub import AudioSegment
from core.utils.util import parse_string_to_list
from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
@@ -13,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
TAG = __name__
logger = setup_logging()
class TTSProvider(TTSProviderBase):
def __init__(self, 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.prompt_text = config.get("prompt_text")
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.batch_size = config.get("batch_size", 1)
self.batch_threshold = config.get("batch_threshold", 0.75)
self.split_bucket = config.get("split_bucket", True)
self.return_fragment = config.get("return_fragment", False)
self.speed_factor = config.get("speed_factor", 1.0)
self.streaming_mode = config.get("streaming_mode", False)
self.seed = config.get("seed", -1)
self.parallel_infer = config.get("parallel_infer", True)
self.repetition_penalty = config.get("repetition_penalty", 1.35)
self.aux_ref_audio_paths = config.get("aux_ref_audio_paths", [])
self.split_bucket = str(config.get("split_bucket", True)).lower() in (
"true",
"1",
"yes",
)
self.return_fragment = str(config.get("return_fragment", False)).lower() in (
"true",
"1",
"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"):
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):
tmp_file = self.generate_filename()
@@ -60,7 +98,7 @@ class TTSProvider(TTSProviderBase):
"streaming_mode": self.streaming_mode,
"seed": self.seed,
"parallel_infer": self.parallel_infer,
"repetition_penalty": self.repetition_penalty
"repetition_penalty": self.repetition_penalty,
}
resp = requests.post(self.url, json=request_json)
@@ -68,13 +106,20 @@ class TTSProvider(TTSProviderBase):
with open(tmp_file, "wb") as file:
file.write(resp.content)
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 读取临时文件
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -2,7 +2,7 @@ import os
import uuid
import requests
from pydub import AudioSegment
from core.utils.util import parse_string_to_list
from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
@@ -11,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
TAG = __name__
logger = setup_logging()
class TTSProvider(TTSProviderBase):
def __init__(self, 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_language = config.get("prompt_language")
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"):
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):
tmp_file = self.generate_filename()
@@ -55,13 +67,20 @@ class TTSProvider(TTSProviderBase):
with open(tmp_file, "wb") as file:
file.write(resp.content)
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 读取临时文件
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -2,10 +2,9 @@ import os
import uuid
import json
import requests
from core.utils.util import parse_string_to_list
from datetime import datetime
from pydub import AudioSegment
from core.providers.tts.base import TTSProviderBase
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.api_key = config.get("api_key")
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 = {
"voice_id": "female-shaonv",
"speed": 1,
"vol": 1,
"pitch": 0,
"emotion": "happy"
}
default_pronunciation_dict = {
"tone": [
"处理/(chu3)(li3)", "危险/dangerous"
]
"emotion": "happy",
}
default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]}
defult_audio_setting = {
"sample_rate": 32000,
"bitrate": 128000,
"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.timber_weights = config.get("timber_weights", [])
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
if 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.header = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}"
"Authorization": f"Bearer {self.api_key}",
}
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):
tmp_file = self.generate_filename()
@@ -70,25 +77,34 @@ class TTSProvider(TTSProviderBase):
request_json["voice_setting"]["voice_id"] = ""
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
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.write(bytes.fromhex(data))
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:
raise Exception(f"{__name__} error: {e}")
# 使用 pydub 读取临时文件
audio = AudioSegment.from_file(tmp_file, format="mp3")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
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.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.api_key = config.get("api_key")
self.api_url = config.get("api_url", "https://api.openai.com/v1/audio/speech")
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.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/")
check_model_key("TTS", self.api_key)
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):
tmp_file = self.generate_filename()
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
"Content-Type": "application/json",
}
data = {
"model": self.model,
"input": text,
"voice": self.voice,
"response_format": "wav",
"speed": self.speed
"speed": self.speed,
}
response = requests.post(self.api_url, json=data, headers=headers)
if response.status_code == 200:
with open(tmp_file, "wb") as audio_file:
audio_file.write(response.content)
else:
raise Exception(f"OpenAI TTS请求失败: {response.status_code} - {response.text}")
raise Exception(
f"OpenAI TTS请求失败: {response.status_code} - {response.text}"
)
# 使用 pydub 读取临时文件
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -14,17 +14,23 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.model = config.get("model")
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.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.host = "api.siliconflow.cn"
self.api_url = f"https://{self.host}/v1/audio/speech"
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):
tmp_file = self.generate_filename()
@@ -36,9 +42,11 @@ class TTSProvider(TTSProviderBase):
}
headers = {
"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
file_to_save = open(tmp_file, "wb")
file_to_save.write(data)
@@ -46,8 +54,13 @@ class TTSProvider(TTSProviderBase):
audio = AudioSegment.from_file(tmp_file, format="wav")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
+48 -30
View File
@@ -14,49 +14,63 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType
class TTSProvider(TTSProviderBase):
def __init__(self, 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.voice_id = config.get("voice_id", 1695)
self.url = config.get(
"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.to_lang = config.get("to_lang")
self.volume_change_dB = config.get("volume_change_dB", 0)
self.speed_factor = config.get("speed_factor", 1)
self.stream = config.get("stream", False)
self.volume_change_dB = int(config.get("volume_change_dB", 0))
self.speed_factor = int(config.get("speed_factor", 1))
self.stream = str(config.get("stream", False)).lower() in ("true", "1", "yes")
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.emotion = config.get("emotion", 1)
self.header = {
"Content-Type": "application/json"
}
self.emotion = int(config.get("emotion", 1))
self.header = {"Content-Type": "application/json"}
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):
tmp_file = self.generate_filename()
url = f'{self.url}{self.token}'
url = f"{self.url}{self.token}"
result = "firefly"
payload = json.dumps({
"to_lang": self.to_lang,
"text": text,
"emotion": self.emotion,
"format": self.format,
"volume_change_dB": self.volume_change_dB,
"voice_id": self.voice_id,
"pitch_factor": self.pitch_factor,
"speed_factor": self.speed_factor,
"token": self.token
})
payload = json.dumps(
{
"to_lang": self.to_lang,
"text": text,
"emotion": self.emotion,
"format": self.format,
"volume_change_dB": self.volume_change_dB,
"voice_id": self.voice_id,
"pitch_factor": self.pitch_factor,
"speed_factor": self.speed_factor,
"token": self.token,
}
)
resp = requests.request("POST", url, data=payload)
if resp.status_code != 200:
return
resp_json = resp.json()
try:
result = resp_json['url'] + ':' + str(
resp_json[
'port']) + '/flashsummary/retrieveFileData?stream=True&token=' + self.token + '&voice_audio_path=' + \
resp_json['voice_path']
result = (
resp_json["url"]
+ ":"
+ str(resp_json["port"])
+ "/flashsummary/retrieveFileData?stream=True&token="
+ self.token
+ "&voice_audio_path="
+ resp_json["voice_path"]
)
except Exception as e:
print("error:", e)
@@ -67,8 +81,13 @@ class TTSProvider(TTSProviderBase):
audio = AudioSegment.from_file(tmp_file, format="mp3")
audio = audio.set_channels(1).set_frame_rate(16000)
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,
sentence_type=SentenceType.SENTENCE_START)
yield TTSMessageDTO(
u_id=u_id,
msg_type=MsgType.TTS_TEXT_RESPONSE,
content=opus_datas,
tts_finish_text=text,
sentence_type=SentenceType.SENTENCE_START,
)
# 用完后删除临时文件
try:
os.remove(tmp_file)
@@ -78,4 +97,3 @@ class TTSProvider(TTSProviderBase):
voice_path = resp_json.get("voice_path")
des_path = tmp_file
shutil.move(voice_path, des_path)