LLM和TTS采用适配器模式进行解耦,方便扩展更多服务平台调用

This commit is contained in:
HonestQiao
2025-02-11 23:13:18 +08:00
parent 0105d97412
commit 28798dde5c
10 changed files with 287 additions and 270 deletions
+16
View File
@@ -27,7 +27,9 @@ delete_audio: true
selected_module: selected_module:
ASR: FunASR ASR: FunASR
VAD: SileroVAD VAD: SileroVAD
# 将根据配置名称对应的type调用实际的LLM适配器
LLM: ChatGLMLLM LLM: ChatGLMLLM
# TTS将根据配置名称对应的type调用实际的TTS适配器
TTS: EdgeTTS TTS: EdgeTTS
ASR: ASR:
@@ -42,31 +44,45 @@ VAD:
min_silence_duration_ms: 700 # 如果说话停顿比较长,可以把这个值设置大一些 min_silence_duration_ms: 700 # 如果说话停顿比较长,可以把这个值设置大一些
LLM: LLM:
# 当前支持的type为openai、dify,可自行适配
AliLLM: AliLLM:
# 定义LLM API类型
type: openai
# 可在这里找到你的 api_key https://bailian.console.aliyun.com/?apiKey=1#/api-key # 可在这里找到你的 api_key https://bailian.console.aliyun.com/?apiKey=1#/api-key
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1 base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
model_name: qwen-turbo model_name: qwen-turbo
api_key: 你的阿里云dashscope api key api_key: 你的阿里云dashscope api key
DeepSeekLLM: DeepSeekLLM:
# 定义LLM API类型
type: openai
# 可在这里找到你的api key https://platform.deepseek.com/ # 可在这里找到你的api key https://platform.deepseek.com/
model_name: deepseek-chat model_name: deepseek-chat
url: https://api.deepseek.com url: https://api.deepseek.com
api_key: 你的deepseek api key api_key: 你的deepseek api key
ChatGLMLLM: ChatGLMLLM:
# 定义LLM API类型
type: openai
# 可在这里找到你的api key https://bigmodel.cn/usercenter/proj-mgmt/apikeys # 可在这里找到你的api key https://bigmodel.cn/usercenter/proj-mgmt/apikeys
model_name: glm-4-flash model_name: glm-4-flash
url: https://open.bigmodel.cn/api/paas/v4/ url: https://open.bigmodel.cn/api/paas/v4/
api_key: 你的bigmodel api key api_key: 你的bigmodel api key
DifyLLM: DifyLLM:
# 定义LLM API类型
type: dify
# 建议使用本地部署的dify接口,国内部分区域访问dify公有云接口可能会受限 # 建议使用本地部署的dify接口,国内部分区域访问dify公有云接口可能会受限
# 如果使用DifyLLM,配置文件里prompt(提示词)是无效的,需要在dify控制台设置提示词 # 如果使用DifyLLM,配置文件里prompt(提示词)是无效的,需要在dify控制台设置提示词
base_url: 你的私有化部署的dify接口地址 base_url: 你的私有化部署的dify接口地址
api_key: 你的dify api key api_key: 你的dify api key
TTS: TTS:
# 当前支持的type为edge、doubao,可自行适配
EdgeTTS: EdgeTTS:
# 定义TTS API类型
type: edge
voice: zh-CN-XiaoxiaoNeural voice: zh-CN-XiaoxiaoNeural
output_file: tmp/ output_file: tmp/
DoubaoTTS: DoubaoTTS:
# 定义TTS API类型
type: doubao
# 火山引擎语音合成服务,需要先在火山引擎控制台创建应用并获取appid和access_token # 火山引擎语音合成服务,需要先在火山引擎控制台创建应用并获取appid和access_token
# 山引擎语音一定要购买花钱,起步价30元,就有100并发了。如果用免费的只有2个并发,会经常报tts错误 # 山引擎语音一定要购买花钱,起步价30元,就有100并发了。如果用免费的只有2个并发,会经常报tts错误
# 购买服务后,购买免费的音色后,可能要等半小时左右,才能使用。 # 购买服务后,购买免费的音色后,可能要等半小时左右,才能使用。
+7
View File
@@ -0,0 +1,7 @@
from abc import ABC, abstractmethod
class LLMProviderBase(ABC):
@abstractmethod
def response(self, session_id, dialogue):
"""LLM response generator"""
pass
+36
View File
@@ -0,0 +1,36 @@
import logging
import requests
from core.providers.llm.base import LLMProviderBase
logger = logging.getLogger(__name__)
class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.api_key = config["api_key"]
self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/')
def response(self, session_id, dialogue):
try:
# 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
# 发起流式请求
with requests.post(
f"{self.base_url}/chat-messages",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"query": last_msg["content"],
"response_mode": "streaming",
"user": session_id,
"inputs": {}
},
stream=True
) as r:
for line in r.iter_lines():
if line.startswith(b'data: '):
event = json.loads(line[6:])
if event.get('answer'):
yield event['answer']
except Exception:
yield "【服务响应异常】"
+32
View File
@@ -0,0 +1,32 @@
import logging
import openai
from core.providers.llm.base import LLMProviderBase
logger = logging.getLogger(__name__)
class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.model_name = config.get("model_name")
self.api_key = config.get("api_key")
if 'base_url' in config:
self.base_url = config.get("base_url")
else:
self.base_url = config.get("url")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue):
try:
responses = self.client.chat.completions.create(
model=self.model_name,
messages=dialogue,
stream=True
)
for chunk in responses:
# 检查是否存在有效的choice且content不为空
if chunk.choices and len(chunk.choices) > 0:
delta = chunk.choices[0].delta
content = getattr(delta, 'content', '')
if content: # 仅在content非空时生成
yield content
except Exception as e:
logger.error(f"Error in response generation: {e}")
+89
View File
@@ -0,0 +1,89 @@
import asyncio
import logging
import os
import json
import uuid
import base64
# from datetime import datetime
# import edge_tts
import numpy as np
import opuslib
# import requests
# from core.utils.util import read_config, get_project_dir
from pydub import AudioSegment
from abc import ABC, abstractmethod
logger = logging.getLogger(__name__)
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.delete_audio_file = delete_audio_file
self.output_file = config.get("output_file")
@abstractmethod
def generate_filename(self):
pass
def to_tts(self, text):
tmp_file = self.generate_filename()
try:
max_repeat_time = 5
while not os.path.exists(tmp_file) and max_repeat_time > 0:
asyncio.run(self.text_to_speak(text, tmp_file))
if not os.path.exists(tmp_file):
max_repeat_time = max_repeat_time - 1
logger.error(f"语音生成失败: {text}:{tmp_file},再试{max_repeat_time}")
if max_repeat_time>0:
logger.info(f"语音生成成功: {text}:{tmp_file},重试{5-max_repeat_time}")
return tmp_file
except Exception as e:
logger.info(f"Failed to generate TTS file: {e}")
return None
@abstractmethod
async def text_to_speak(self, text, output_file):
pass
def wav_to_opus_data(self, wav_file_path):
# 使用pydub加载PCM文件
# 获取文件后缀名
file_type = os.path.splitext(wav_file_path)[1]
if file_type:
file_type = file_type.lstrip('.')
audio = AudioSegment.from_file(wav_file_path, format=file_type)
duration = len(audio) / 1000.0
# 转换为单声道和16kHz采样率(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000)
# 获取原始PCM数据(16位小端)
raw_data = audio.raw_data
# 初始化Opus编码器
encoder = opuslib.Encoder(16000, 1, opuslib.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
+55
View File
@@ -0,0 +1,55 @@
import os
import uuid
import json
import base64
import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.appid = config.get("appid")
self.access_token = config.get("access_token")
self.cluster = config.get("cluster")
self.voice = config.get("voice")
self.host = "openspeech.bytedance.com"
self.api_url = f"https://{self.host}/api/v1/tts"
self.header = {"Authorization": f"Bearer;{self.access_token}"}
def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
request_json = {
"app": {
"appid": self.appid,
"token": "access_token",
"cluster": self.cluster
},
"user": {
"uid": "1"
},
"audio": {
"voice_type": self.voice,
"encoding": "wav",
"speed_ratio": 1.0,
"volume_ratio": 1.0,
"pitch_ratio": 1.0,
},
"request": {
"reqid": str(uuid.uuid4()),
"text": text,
"text_type": "plain",
"operation": "query",
"with_frontend": 1,
"frontend_type": "unitTson"
}
}
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header)
if "data" in resp.json():
data = resp.json()["data"]
file_to_save = open(output_file, "wb")
file_to_save.write(base64.b64decode(data))
+17
View File
@@ -0,0 +1,17 @@
import os
import uuid
import edge_tts
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.voice = config.get("voice")
def generate_filename(self, extension=".mp3"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice
await communicate.save(output_file)
+8 -2
View File
@@ -25,11 +25,17 @@ class WebSocketServer:
self.config["delete_audio"] self.config["delete_audio"]
), ),
llm.create_instance( llm.create_instance(
self.config["selected_module"]["LLM"], self.config["selected_module"]["LLM"]
if not 'type' in self.config["LLM"][self.config["selected_module"]["LLM"]]
else
self.config["LLM"][self.config["selected_module"]["LLM"]]['type'],
self.config["LLM"][self.config["selected_module"]["LLM"]], self.config["LLM"][self.config["selected_module"]["LLM"]],
), ),
tts.create_instance( tts.create_instance(
self.config["selected_module"]["TTS"], self.config["selected_module"]["TTS"]
if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]]
else
self.config["TTS"][self.config["selected_module"]["TTS"]]["type"],
self.config["TTS"][self.config["selected_module"]["TTS"]], self.config["TTS"][self.config["selected_module"]["TTS"]],
self.config["delete_audio"] self.config["delete_audio"]
) )
+14 -125
View File
@@ -1,7 +1,10 @@
import os
import sys
import json import json
import logging import logging
import openai import openai
import requests import requests
import importlib
from datetime import datetime from datetime import datetime
from core.utils.util import is_segment from core.utils.util import is_segment
from core.utils.util import get_string_no_punctuation_or_emoji from core.utils.util import get_string_no_punctuation_or_emoji
@@ -11,132 +14,15 @@ from abc import ABC, abstractmethod
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class LLM(ABC):
@abstractmethod
def response(self, session_id, dialogue):
"""LLM response generator"""
pass
class DeepSeekLLM(LLM):
def __init__(self, config):
self.model_name = config.get("model_name")
self.api_key = config.get("api_key")
self.base_url = config.get("url")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue):
logger.info(f"Generating response using {dialogue}")
try:
responses = self.client.chat.completions.create(
model=self.model_name,
messages=dialogue,
stream=True
)
for chunk in responses:
# 检查是否存在有效的choice且content不为空
if chunk.choices and len(chunk.choices) > 0:
delta = chunk.choices[0].delta
content = getattr(delta, 'content', '')
if content: # 仅在content非空时生成
yield content
except Exception as e:
logger.error(f"Error in response generation: {e}")
class ChatGLMLLM(LLM):
def __init__(self, config):
self.model_name = config.get("model_name")
self.api_key = config.get("api_key")
self.base_url = config.get("url")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue):
try:
responses = self.client.chat.completions.create(
model=self.model_name,
messages=dialogue,
stream=True
)
for chunk in responses:
# 检查是否存在有效的choice且content不为空
if chunk.choices and len(chunk.choices) > 0:
delta = chunk.choices[0].delta
content = getattr(delta, 'content', '')
if content: # 仅在content非空时生成
yield content
except Exception as e:
logger.error(f"Error in response generation: {e}")
class AliLLM(LLM):
def __init__(self, config):
self.model_name = config.get("model_name")
self.api_key = config.get("api_key")
self.base_url = config.get("base_url")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue):
try:
responses = self.client.chat.completions.create(
model=self.model_name,
messages=dialogue,
stream=True
)
for chunk in responses:
# 检查是否存在有效的choice且content不为空
if chunk.choices and len(chunk.choices) > 0:
delta = chunk.choices[0].delta
content = getattr(delta, 'content', '')
if content: # 仅在content非空时生成
yield content
except Exception as e:
logger.error(f"Error in response generation: {e}")
class DifyLLM(LLM):
def __init__(self, config):
self.api_key = config["api_key"]
self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/')
def response(self, session_id, dialogue):
try:
# 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
# 发起流式请求
with requests.post(
f"{self.base_url}/chat-messages",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"query": last_msg["content"],
"response_mode": "streaming",
"user": session_id,
"inputs": {}
},
stream=True
) as r:
for line in r.iter_lines():
if line.startswith(b'data: '):
event = json.loads(line[6:])
if event.get('answer'):
yield event['answer']
except Exception:
yield "【服务响应异常】"
def create_instance(class_name, *args, **kwargs): def create_instance(class_name, *args, **kwargs):
# 获取类对象 # 创建LLM实例
cls_map = { if os.path.exists(os.path.join('core', 'providers', 'llm', class_name, f'{class_name}.py')):
"DeepSeekLLM": DeepSeekLLM, lib_name = f'core.providers.llm.{class_name}.{class_name}'
"ChatGLMLLM": ChatGLMLLM, if lib_name not in sys.modules:
"DifyLLM": DifyLLM, sys.modules[lib_name] = importlib.import_module(f'{lib_name}')
"AliLLM": AliLLM, return sys.modules[lib_name].LLMProvider(*args, **kwargs)
# 可扩展其他LLM实现
}
if cls := cls_map.get(class_name): raise ValueError(f"不支持的LLM类型: {class_name},请检查该配置的type是否设置正确")
return cls(*args, **kwargs)
raise ValueError(f"不支持的LLM类型: {class_name}")
if __name__ == "__main__": if __name__ == "__main__":
@@ -145,7 +31,10 @@ if __name__ == "__main__":
""" """
config = read_config(get_project_dir() + "config.yaml") config = read_config(get_project_dir() + "config.yaml")
llm = create_instance( llm = create_instance(
config["selected_module"]["LLM"], config["selected_module"]["LLM"]
if not "type" in config["LLM"][config["selected_module"]["LLM"]]
else
config["LLM"][config["selected_module"]["LLM"]]["type"],
config["LLM"][config["selected_module"]["LLM"]] config["LLM"][config["selected_module"]["LLM"]]
) )
+13 -143
View File
@@ -1,9 +1,11 @@
import asyncio import asyncio
import logging import logging
import os import os
import sys
import json import json
import uuid import uuid
import base64 import base64
import importlib
from datetime import datetime from datetime import datetime
import edge_tts import edge_tts
import numpy as np import numpy as np
@@ -16,150 +18,15 @@ from abc import ABC, abstractmethod
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class TTS(ABC):
def __init__(self, config, delete_audio_file):
self.delete_audio_file = delete_audio_file
self.output_file = config.get("output_file")
@abstractmethod
def generate_filename(self):
pass
def to_tts(self, text):
tmp_file = self.generate_filename()
try:
max_repeat_time = 5
while not os.path.exists(tmp_file) and max_repeat_time > 0:
asyncio.run(self.text_to_speak(text, tmp_file))
if not os.path.exists(tmp_file):
max_repeat_time = max_repeat_time - 1
logger.error(f"语音生成失败: {text}:{tmp_file},再试{max_repeat_time}")
return tmp_file
except Exception as e:
logger.info(f"Failed to generate TTS file: {e}")
return None
@abstractmethod
async def text_to_speak(self, text, output_file):
pass
def wav_to_opus_data(self, wav_file_path):
# 使用pydub加载PCM文件
# 获取文件后缀名
file_type = os.path.splitext(wav_file_path)[1]
if file_type:
file_type = file_type.lstrip('.')
audio = AudioSegment.from_file(wav_file_path, format=file_type)
duration = len(audio) / 1000.0
# 转换为单声道和16kHz采样率(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000)
# 获取原始PCM数据(16位小端)
raw_data = audio.raw_data
# 初始化Opus编码器
encoder = opuslib.Encoder(16000, 1, opuslib.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
class EdgeTTS(TTS):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.voice = config.get("voice")
def generate_filename(self, extension=".mp3"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice
await communicate.save(output_file)
class DoubaoTTS(TTS):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.appid = config.get("appid")
self.access_token = config.get("access_token")
self.cluster = config.get("cluster")
self.voice = config.get("voice")
self.host = "openspeech.bytedance.com"
self.api_url = f"https://{self.host}/api/v1/tts"
self.header = {"Authorization": f"Bearer;{self.access_token}"}
def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
async def text_to_speak(self, text, output_file):
request_json = {
"app": {
"appid": self.appid,
"token": "access_token",
"cluster": self.cluster
},
"user": {
"uid": "1"
},
"audio": {
"voice_type": self.voice,
"encoding": "wav",
"speed_ratio": 1.0,
"volume_ratio": 1.0,
"pitch_ratio": 1.0,
},
"request": {
"reqid": str(uuid.uuid4()),
"text": text,
"text_type": "plain",
"operation": "query",
"with_frontend": 1,
"frontend_type": "unitTson"
}
}
resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header)
if "data" in resp.json():
data = resp.json()["data"]
file_to_save = open(output_file, "wb")
file_to_save.write(base64.b64decode(data))
def create_instance(class_name, *args, **kwargs): def create_instance(class_name, *args, **kwargs):
# 获取类对象 # 创建TTS实例
cls_map = { if os.path.exists(os.path.join('core', 'providers', 'tts', f'{class_name}.py')):
"DoubaoTTS": DoubaoTTS, lib_name = f'core.providers.tts.{class_name}'
"EdgeTTS": EdgeTTS, if lib_name not in sys.modules:
# 可扩展其他TTS实现 sys.modules[lib_name] = importlib.import_module(f'{lib_name}')
} return sys.modules[lib_name].TTSProvider(*args, **kwargs)
if cls := cls_map.get(class_name): raise ValueError(f"不支持的TTS类型: {class_name},请检查该配置的type是否设置正确")
return cls(*args, **kwargs)
raise ValueError(f"不支持的TTS类型: {class_name}")
if __name__ == "__main__": if __name__ == "__main__":
@@ -168,7 +35,10 @@ if __name__ == "__main__":
""" """
config = read_config(get_project_dir() + "config.yaml") config = read_config(get_project_dir() + "config.yaml")
tts = create_instance( tts = create_instance(
config["selected_module"]["TTS"], config["selected_module"]["TTS"]
if not 'type' in config["TTS"][config["selected_module"]["TTS"]]
else
config["TTS"][config["selected_module"]["TTS"]]["type"],
config["TTS"][config["selected_module"]["TTS"]], config["TTS"][config["selected_module"]["TTS"]],
config["delete_audio"] config["delete_audio"]
) )