Compare commits

...
14 Commits
Author SHA1 Message Date
87cb0b484d Music play (#142)
* Update util.py解决返回元组报错 (#103)

解决返回元组报错

* 增加播放本地音乐功能 (#105)

* Please enter the commit message for your changes. Lines starting
 with '#' will be ignored, and an empty message aborts the commit.

 On branch Music-playback
 Your branch is up to date with 'origin/Music-playback'.

 Changes to be committed:
	modified:   core/handle/audioHandle.py
	new file:   "music/\346\234\210\344\272\256\344\273\243\350\241\250\346\210\221\347\232\204\345\277\203_\351\202\223\344\270\275\345\220\233.mp3"
	new file:   "music/\350\270\217\345\261\261\346\262\263.mp3"

* Please enter the commit message for your changes. Lines starting
with '#' will be ignored, and an empty message aborts the commit.

On branch Music-playback
Your branch is up to date with 'origin/Music-playback'.

Changes to be committed:
	modified:   config.yaml
	modified:   core/connection.py
	modified:   core/handle/abortHandle.py
	modified:   core/handle/audioHandle.py
	new file:   core/handle/musicHandler.py
	modified:   core/handle/textHandle.py
	modified:   core/websocket_server.py
	new file:   "music/\344\270\200\345\277\265\345\215\203\345\271\264_\345\233\275\351\243\216\347\211\210.mp3"
	new file:   "music/\344\270\255\347\247\213\346\234\210.mp3"
	deleted:    "music/\346\234\210\344\272\256\344\273\243\350\241\250\346\210\221\347\232\204\345\277\203_\351\202\223\344\270\275\345\220\233.mp3"
	deleted:    "music/\350\270\217\345\261\261\346\262\263.mp3"

* 增加播放本地音乐功能

* 增加播放本地音乐功能

* 让音乐配置变的优雅

On branch Music-playback
Your branch is up to date with 'origin/Music-playback'.

Changes to be committed:
	modified:   config.yaml
	modified:   core/handle/musicHandler.py
	renamed:    "music/\344\270\200\345\277\265\345\215\203\345\271\264_\345\233\275\351\243\216\347\211\210.mp3" -> "music/mp3/\344\270\200\345\277\265\345\215\203\345\271\264_\345\233\275\351\243\216\347\211\210.mp3"
	renamed:    "music/\344\270\255\347\247\213\346\234\210.mp3" -> "music/mp3/\344\270\255\347\247\213\346\234\210.mp3"
	renamed:    "music/\345\273\211\346\263\242\350\200\201\347\237\243\357\274\214\345\260\232\350\203\275\351\245\255\345\220\246.mp3" -> "music/mp3/\345\273\211\346\263\242\350\200\201\347\237\243\357\274\214\345\260\232\350\203\275\351\245\255\345\220\246.mp3"
	new file:   music/music_config.yaml

---------

Co-authored-by: 欣南科技 <huangrongzhuang@xin-nan.com>

* update:优化代码

* update:流控调试

* update:优化音乐播放

* update:优化音乐配置初始化

---------

Co-authored-by: linqingping <linqingping@users.noreply.github.com>
Co-authored-by: Chris <119588753+Chris-websketch@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-26 01:33:05 +08:00
73803f49fc update:wechat group (#141)
Co-authored-by: hrz <1710360675@qq.com>
2025-02-25 23:51:57 +08:00
5fd87752d5 fixed:更新基础docker环境地址 (#139)
Co-authored-by: hrz <1710360675@qq.com>
2025-02-25 21:11:06 +08:00
94a90c4124 Iot (#129)
* 2025-2-17-iot配置消息处理, 以及控制iot命令发送 (#40)

* update:优化

* update:增加视频教程

* update:iot兼容性又问题,撤回

* fixed:iot boolean 参数接收bug

---------

Co-authored-by: Jiao Haoyang <108573524+XuSenfeng@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 22:04:48 +08:00
c3b0347155 Iot (#128)
* 2025-2-17-iot配置消息处理, 以及控制iot命令发送 (#40)

* update:优化

* update:增加视频教程

* update:iot兼容性又问题,撤回

---------

Co-authored-by: Jiao Haoyang <108573524+XuSenfeng@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 21:20:42 +08:00
600f712750 Iot (#127)
* 2025-2-17-iot配置消息处理, 以及控制iot命令发送 (#40)

* update:优化

* update:增加视频教程

---------

Co-authored-by: Jiao Haoyang <108573524+XuSenfeng@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 18:18:08 +08:00
9e76869c36 Iot (#126)
* 2025-2-17-iot配置消息处理, 以及控制iot命令发送 (#40)

* update:优化

---------

Co-authored-by: Jiao Haoyang <108573524+XuSenfeng@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 18:06:13 +08:00
cb540736ab 跳过使用Open ai 接口时,DeepSeek-R1 模型的深度思考内容 (#54) (#95)
* 跳过使用Open ai 接口时,DeepSeek-R1 模型的深度思考内容 (#54)

* 跳过 DeepSeek-R1 模型的深度思考内容

* 跳过 DeepSeek-R1 模型的深度思考内容

* 新增LM Studio本地大模型API接口

* 优化代码,遇到Bad Case安全处理

* update:优化

---------

Co-authored-by: Sinyo <38577585+SinyoWong@users.noreply.github.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 16:16:06 +08:00
66cd5bb4a5 fixed:修复低版本docker-compose模版参数 (#124)
Co-authored-by: hrz <1710360675@qq.com>
2025-02-24 14:56:39 +08:00
1226e8be47 Fix opuslib (#122)
* fixed:ModuleNotFoundError: No module named 'opuslib'

* update:格式化文档

---------

Co-authored-by: hrz <1710360675@qq.com>
2025-02-23 20:14:46 +08:00
22ff7a92b8 fixed:ModuleNotFoundError: No module named 'opuslib' (#121)
Co-authored-by: hrz <1710360675@qq.com>
2025-02-23 19:59:26 +08:00
2d01812f8d Performance test (#118)
* update: optimize performance test files

* update:异步测试方法

* update:改回默认FunASR

---------

Co-authored-by: Alexisxty <alexisty233@gmail.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-23 16:22:21 +08:00
3bc0b821a1 update: 修改了服务组件的性能测试代码 (#111)
* update: optimize performance test files

* update:异步测试方法

---------

Co-authored-by: hrz <1710360675@qq.com>
2025-02-23 16:07:33 +08:00
148578399f 使用配置文件设置日志参数 (#116)
* 使用配置文件设置日志参数

* update:优化

---------

Co-authored-by: 欣南科技 <huangrongzhuang@xin-nan.com>
Co-authored-by: hrz <1710360675@qq.com>
2025-02-23 15:18:59 +08:00
22 changed files with 928 additions and 362 deletions
+19 -33
View File
@@ -58,28 +58,13 @@
</tr> </tr>
</table> </table>
--- ---
## 系统要求与部署前提 🖥️ ## 系统要求与部署前提 🖥️
- **硬件**:一套兼容 `xiaozhi-esp32` - **硬件**:一套兼容 `xiaozhi-esp32`
的硬件设备(具体型号请参考 [此处](https://rcnv1t9vps13.feishu.cn/wiki/DdgIw4BUgivWDPkhMj1cGIYCnRf))。 的硬件设备(具体型号请参考 [此处](https://rcnv1t9vps13.feishu.cn/wiki/DdgIw4BUgivWDPkhMj1cGIYCnRf))。
- **服务器**:至少 4 核 CPU、8G 内存的电脑或服务器 - **服务器**:至少 4 核 CPU、8G 内存的电脑。
- **固件编译**:请将后端服务的接口地址更新至 `xiaozhi-esp32` 项目中,再重新编译固件并烧录到设备上。 - **固件编译**:请将后端服务的接口地址更新至 `xiaozhi-esp32` 项目中,再重新编译固件并烧录到设备上。
--- ---
@@ -167,9 +152,10 @@ server:
### ASR ### ASR
| 类型 | 平台名称 | 使用方式 | 收费模式 | 备注 | | 类型 | 平台名称 | 使用方式 | 收费模式 | 备注 |
|:---:|:------:|:----:|:----:|:--:| |:---:|:---------:|:----:|:----:|:--:|
| ASR | FunASR | 本地使用 | 免费 | | | ASR | FunASR | 本地使用 | 免费 | |
| ASR | DoubaoASR | 接口调用 | 收费 | |
--- ---
@@ -177,23 +163,24 @@ server:
### 一、[部署文档](./docs/Deployment.md) ### 一、[部署文档](./docs/Deployment.md)
本项目支持以下三种部署方式,您可根据实际需求选择 本项目支持以下三种部署方式,您可根据实际需求选择
1. **[Docker 快速部署](./docs/Deployment.md)** 本项目的文档主要是`文字版本`的教程,如果你想要`视频版本`的教程,您可以学习一下[这个大佬的手把手教程](https://www.bilibili.com/video/BV1gePuejEvT)。
适合快速体验,不需过多环境配置。缺点是,拉取镜像有点慢。
2.
*
*[借助 Docker 环境运行部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E5%80%9F%E5%8A%A9docker%E7%8E%AF%E5%A2%83%E8%BF%90%E8%A1%8C%E9%83%A8%E7%BD%B2) 如果你能把`文字版本的教程``视频版本的教程`结合起来一起看,可以让你更快上手。
**
适用于已安装 Docker 且希望对代码进行自定义修改的用户。
3. 1. [Docker 快速部署](./docs/Deployment.md)
*
适合快速体验的普通用户,不需过多环境配置。缺点是,拉取镜像有点慢。
2. [借助 Docker 环境运行部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E5%80%9F%E5%8A%A9docker%E7%8E%AF%E5%A2%83%E8%BF%90%E8%A1%8C%E9%83%A8%E7%BD%B2)
适用于已安装 Docker 且希望对代码进行自定义修改的软件工程师。
3. [本地源码运行](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%89%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C)
适合熟悉`Conda` 环境或希望从零搭建运行环境的用户。
*[本地源码运行](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%89%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C)
**
适合熟悉 Conda 环境或希望从零搭建运行环境的用户。
对于对响应速度要求较高的场景,推荐使用本地源码运行方式以降低额外开销。 对于对响应速度要求较高的场景,推荐使用本地源码运行方式以降低额外开销。
### 二、[固件编译](./docs/firmware-build.md) ### 二、[固件编译](./docs/firmware-build.md)
@@ -310,7 +297,6 @@ TTS 性能排行:
## 鸣谢 🙏 ## 鸣谢 🙏
- 本项目受 [百聆语音对话机器人](https://github.com/wwbin2017/bailing) 启发,并在其基础上实现。 - 本项目受 [百聆语音对话机器人](https://github.com/wwbin2017/bailing) 启发,并在其基础上实现。
- 感谢 [腾讯云](https://cloud.tencent.com/) 提供免费 Docker 镜像空间。
- 感谢 [十方融海](https://www.tenclass.com/) 对小智通讯协议提供的详尽文档支持。 - 感谢 [十方融海](https://www.tenclass.com/) 对小智通讯协议提供的详尽文档支持。
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date"> <a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
+39 -2
View File
@@ -22,13 +22,28 @@ server:
# 可选:设备白名单,如果设置了白名单,那么白名单的机器无论是什么token都可以连接。 # 可选:设备白名单,如果设置了白名单,那么白名单的机器无论是什么token都可以连接。
#allowed_devices: #allowed_devices:
# - "24:0A:C4:1D:3B:F0" # MAC地址列表 # - "24:0A:C4:1D:3B:F0" # MAC地址列表
log:
# 设置控制台输出的日志格式,时间、日志级别、标签、消息
log_format: "<green>{time:YY-MM-DD HH:mm:ss}</green>[<light-blue>{extra[tag]}</light-blue>] - <level>{level}</level> - <light-green>{message}</light-green>"
# 设置日志文件输出的格式,时间、日志级别、标签、消息
log_format_simple: "{time:YYYY-MM-DD HH:mm:ss} - {name} - {level} - {extra[tag]} - {message}"
# 设置日志等级:INFO、DEBUG
log_level: INFO
# 设置日志路径
log_dir: tmp
# 设置日志文件
log_file: "server.log"
# 设置数据文件路径
data_dir: data
manager: manager:
# 是否启用管理后台 # 是否启用管理后台
# 目前这个模块还在开发中,建议:不要修改enabled选项 # 目前这个模块还在开发中,建议:不要修改enabled选项
enabled: false enabled: false
ip: 0.0.0.0 ip: 0.0.0.0
port: 8002 port: 8002
iot:
Speaker:
volume: 100
xiaozhi: xiaozhi:
type: hello type: hello
version: 1 version: 1
@@ -59,7 +74,7 @@ CMD_exit:
# 具体处理时选择的模块(The module selected for specific processing) # 具体处理时选择的模块(The module selected for specific processing)
selected_module: selected_module:
ASR: DoubaoASR ASR: FunASR
VAD: SileroVAD VAD: SileroVAD
# 将根据配置名称对应的type调用实际的LLM适配器 # 将根据配置名称对应的type调用实际的LLM适配器
LLM: ChatGLMLLM LLM: ChatGLMLLM
@@ -134,6 +149,12 @@ LLM:
user_id: 你的user_id user_id: 你的user_id
base_url: "https://api.coze.cn/open_api/v2/chat" # 服务地址 base_url: "https://api.coze.cn/open_api/v2/chat" # 服务地址
personal_access_token: 你的coze个人令牌 personal_access_token: 你的coze个人令牌
LMStudioLLM:
# 定义LLM API类型
type: openai
model_name: deepseek-r1-distill-llama-8b@q4_k_m # 使用的模型名称,需要预先在社区下载
url: http://localhost:1234/v1 # LM Studio服务地址
api_key: lm-studio # LM Studio服务的固定API Key
HomeAssistant: HomeAssistant:
# 定义LLM API类型 # 定义LLM API类型
type: homeassistant type: homeassistant
@@ -291,3 +312,19 @@ module_test:
- "你好,请介绍一下你自己" - "你好,请介绍一下你自己"
- "What's the weather like today?" - "What's the weather like today?"
- "请用100字概括量子计算的基本原理和应用前景" - "请用100字概括量子计算的基本原理和应用前景"
# 本地音乐播放配置
music:
music_commands:
- "来一首歌"
- "唱一首歌"
- "播放音乐"
- "来点音乐"
- "背景音乐"
- "放首歌"
- "播放歌曲"
- "来点背景音乐"
- "我想听歌"
- "我要听歌"
- "放点音乐"
music_dir: "./music" # 音乐文件存放路径
+14 -12
View File
@@ -1,27 +1,29 @@
import os import os
import sys import sys
from loguru import logger from loguru import logger
from config.settings import load_config
def setup_logging():
"""从配置文件中读取日志配置,并设置日志输出格式和级别"""
config = load_config()
log_config = config["log"]
log_format = log_config.get("log_format", "<green>{time:YY-MM-DD HH:mm:ss}</green>[<light-blue>{extra[tag]}</light-blue>] - <level>{level}</level> - <light-green>{message}</light-green>")
log_format_simple = log_config.get("log_format_file", "{time:YYYY-MM-DD HH:mm:ss} - {name} - {level} - {extra[tag]} - {message}")
log_level = log_config.get("log_level", "INFO")
log_dir = log_config.get("log_dir", "tmp")
log_file = log_config.get("log_file", "server.log")
data_dir = log_config.get("data_dir", "data")
def setup_logging(log_dir='tmp', data_dir='data'):
"""配置全局彩色日志(不同区块不同标签)"""
os.makedirs(log_dir, exist_ok=True) os.makedirs(log_dir, exist_ok=True)
os.makedirs(data_dir, exist_ok=True) os.makedirs(data_dir, exist_ok=True)
# 设置日志格式,时间、日志级别、标签、消息
log_format = (
"<green>{time:YY-MM-DD HH:mm:ss}</green>"
"[<light-blue>{extra[tag]}</light-blue>]"
" - <level>{level}</level> - "
"<light-green>{message}</light-green>"
)
# 配置日志输出 # 配置日志输出
logger.remove() logger.remove()
# 输出到控制台 # 输出到控制台
logger.add(sys.stdout, format=log_format, level="INFO") logger.add(sys.stdout, format=log_format, level=log_level)
# 输出到文件 # 输出到文件
logger.add(os.path.join(log_dir, "server.log"), format="{time:YYYY-MM-DD HH:mm:ss} - {name} - {level} - {extra[tag]} - {message}", level="INFO") logger.add(os.path.join(log_dir, log_file), format=log_format_simple, level=log_level)
return logger return logger
+26 -14
View File
@@ -4,6 +4,7 @@ import uuid
import time import time
import queue import queue
import asyncio import asyncio
import traceback
from config.logger import setup_logging from config.logger import setup_logging
import threading import threading
import websockets import websockets
@@ -14,15 +15,18 @@ from core.utils.dialogue import Message, Dialogue
from core.handle.textHandle import handleTextMessage from core.handle.textHandle import handleTextMessage
from core.utils.util import get_string_no_punctuation_or_emoji from core.utils.util import get_string_no_punctuation_or_emoji
from concurrent.futures import ThreadPoolExecutor, TimeoutError from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.audioHandle import handleAudioMessage, sendAudioMessage from core.handle.sendAudioHandle import sendAudioMessage
from core.handle.receiveAudioHandle import handleAudioMessage
from config.private_config import PrivateConfig from config.private_config import PrivateConfig
from core.auth import AuthMiddleware, AuthenticationError from core.auth import AuthMiddleware, AuthenticationError
from core.utils.auth_code_gen import AuthCodeGenerator # 添加导入 from core.utils.auth_code_gen import AuthCodeGenerator
TAG = __name__ TAG = __name__
class ConnectionHandler: class ConnectionHandler:
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts): def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music):
self.config = config self.config = config
self.logger = setup_logging() self.logger = setup_logging()
self.auth = AuthMiddleware(config) self.auth = AuthMiddleware(config)
@@ -72,28 +76,33 @@ class ConnectionHandler:
self.tts_start_speak_time = None self.tts_start_speak_time = None
self.tts_duration = 0 self.tts_duration = 0
# iot相关变量
self.iot_descriptors = {}
self.cmd_exit = self.config["CMD_exit"] self.cmd_exit = self.config["CMD_exit"]
self.max_cmd_length = 0 self.max_cmd_length = 0
for cmd in self.cmd_exit: for cmd in self.cmd_exit:
if len(cmd) > self.max_cmd_length: if len(cmd) > self.max_cmd_length:
self.max_cmd_length = len(cmd) self.max_cmd_length = len(cmd)
self.private_config = None self.private_config = None
self.auth_code_gen = AuthCodeGenerator.get_instance() self.auth_code_gen = AuthCodeGenerator.get_instance()
self.is_device_verified = False # 添加设备验证状态标志 self.is_device_verified = False # 添加设备验证状态标志
self.music_handler = _music
async def handle_connection(self, ws): async def handle_connection(self, ws):
try: try:
# 获取并验证headers # 获取并验证headers
self.headers = dict(ws.request.headers) self.headers = dict(ws.request.headers)
self.logger.bind(tag=TAG).info(f"New connection request - Headers: {self.headers}") # 获取客户端ip地址
client_ip = ws.remote_address[0]
self.logger.bind(tag=TAG).info(f"{client_ip} conn - Headers: {self.headers}")
# 进行认证 # 进行认证
await self.auth.authenticate(self.headers) await self.auth.authenticate(self.headers)
device_id = self.headers.get("device-id", None) device_id = self.headers.get("device-id", None)
# Load private configuration if device_id is provided # Load private configuration if device_id is provided
bUsePrivateConfig = self.config.get("use_private_config", False) bUsePrivateConfig = self.config.get("use_private_config", False)
self.logger.bind(tag=TAG).info(f"bUsePrivateConfig: {bUsePrivateConfig}, device_id: {device_id}") self.logger.bind(tag=TAG).info(f"bUsePrivateConfig: {bUsePrivateConfig}, device_id: {device_id}")
@@ -104,10 +113,10 @@ class ConnectionHandler:
# 判断是否已经绑定 # 判断是否已经绑定
owner = self.private_config.get_owner() owner = self.private_config.get_owner()
self.is_device_verified = owner is not None self.is_device_verified = owner is not None
if self.is_device_verified: if self.is_device_verified:
await self.private_config.update_last_chat_time() await self.private_config.update_last_chat_time()
llm, tts = self.private_config.create_private_instances() llm, tts = self.private_config.create_private_instances()
if all([llm, tts]): if all([llm, tts]):
self.llm = llm self.llm = llm
@@ -146,7 +155,8 @@ class ConnectionHandler:
await ws.close() await ws.close()
return return
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}") stack_trace = traceback.format_exc()
self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}-{stack_trace}")
await ws.close() await ws.close()
return return
@@ -166,7 +176,7 @@ class ConnectionHandler:
date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime()) date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime())
self.prompt = self.prompt.replace("{date_time}", date_time) self.prompt = self.prompt.replace("{date_time}", date_time)
self.dialogue.put(Message(role="system", content=self.prompt)) self.dialogue.put(Message(role="system", content=self.prompt))
async def _check_and_broadcast_auth_code(self): async def _check_and_broadcast_auth_code(self):
"""检查设备绑定状态并广播认证码""" """检查设备绑定状态并广播认证码"""
if not self.private_config.get_owner(): if not self.private_config.get_owner():
@@ -186,7 +196,7 @@ class ConnectionHandler:
# 如果不使用私有配置,就不需要验证 # 如果不使用私有配置,就不需要验证
return False return False
return not self.is_device_verified return not self.is_device_verified
def chat(self, query): def chat(self, query):
# 如果设备未验证,就发送验证码 # 如果设备未验证,就发送验证码
if self.isNeedAuth(): if self.isNeedAuth():
@@ -199,7 +209,7 @@ class ConnectionHandler:
finally: finally:
loop.close() loop.close()
return True return True
self.dialogue.put(Message(role="user", content=query)) self.dialogue.put(Message(role="user", content=query))
response_message = [] response_message = []
start = 0 start = 0
@@ -316,6 +326,8 @@ class ConnectionHandler:
async def close(self): async def close(self):
"""资源清理方法""" """资源清理方法"""
# 清理其他资源
self.stop_event.set() self.stop_event.set()
self.executor.shutdown(wait=False) self.executor.shutdown(wait=False)
if self.websocket: if self.websocket:
+152
View File
@@ -0,0 +1,152 @@
import json
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class IotDescriptor:
"""
A class to represent an IoT descriptor.
Attributes:
----------
name : str
The name of the IoT descriptor.
description : str
A brief description of the IoT descriptor.
properties : dict
A dictionary containing properties of the IoT descriptor.
methods : dict
A dictionary containing methods of the IoT descriptor.
-------
"""
def __init__(self, name, description, properties, methods):
self.name = name
self.description = description
self.properties = []
self.methods = []
# 根据描述创建属性
for key, value in properties.items():
# "volume":{"description":"当前音量 值","type":"number"}
"""
等价于
{
'name': 名字,
'description': 描述,
'value': 0
}
"""
# setattr(self, key, {}) # 创建一个空字典, 名字是属性名
property_item = globals()[key] = {} # 创建一个空字典, 名字是属性名
property_item['name'] = key
property_item["description"] = value["description"]
if value["type"] == "number":
property_item["value"] = 0
elif value["type"] == "boolean":
property_item["value"] = False
else:
property_item["value"] = ""
self.properties.append(property_item)
# 根据描述创建方法
for key, value in methods.items():
# "SetVolume": {"description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}}}
"""
等价于
SetVolume = {
`description`: 描述,
`volume`: {
`description`: 描述,
`value`: 0
}
}
"""
# setattr(self, key, {}) # 创建一个空字典, 名字是方法名
method = globals()[key] = {} # 创建一个空字典, 名字是方法名
method["description"] = value["description"]
method['name'] = key
for k, v in value["parameters"].items():
# 不同的参数解析
method[k] = {}
method[k]["description"] = v["description"]
if v["type"] == "number":
method[k]["value"] = 0
elif v["type"] == "boolean":
method[k]["value"] = False
else:
method[k]["value"] = ""
self.methods.append(method)
async def handleIotDescriptors(conn, descriptors):
"""
处理物联网描述
示例: [{
"name":"Speaker",
"description":"当前 AI 机器人的扬声器",
"properties":{
"volume":{"description":"当前音量 值","type":"number"} 可以有boolean, number, string三种类型
},
"methods":{
"SetVolume":{
"description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}}
}
}
}]
descriptors: 描述列表
"""
for descriptor in descriptors:
iot_descriptor = IotDescriptor(descriptor["name"], descriptor["description"], descriptor["properties"],
descriptor["methods"])
conn.iot_descriptors[descriptor["name"]] = iot_descriptor
# 暂时从配置文件中设置音量,后期通过意图识别控制音量
default_iot_volume = 100
if "iot" in conn.config:
default_iot_volume = conn.config["iot"]["Speaker"]["volume"]
logger.bind(tag=TAG).info(f"服务端设置音量为{default_iot_volume}")
await send_iot_conn(conn, "Speaker", "SetVolume", {"volume": default_iot_volume})
async def send_iot_conn(conn, name, method_name, parameters):
"""
发送物联网指令
name: 设备名称 "Speaker"
method: 方法 "SetVolume"
parameters: 参数, 是一个字典 {"volume": 100}
发送示例:
{
"type": "iot",
"commands": [
{
"name" : "Speaker",
"method": "SetVolume",
"parameters": {
"volume": 100
}
}
]
}
"""
for key, value in conn.iot_descriptors.items():
if key == name:
# 找到了设备
for method in value.methods:
# 找到了方法
if method["name"] == method_name:
await conn.websocket.send(json.dumps({
"type": "iot",
"commands": [
{
"name": name,
"method": method_name,
"parameters": parameters
}
]
}))
return
logger.bind(tag=TAG).error(f"未找到方法{method_name}")
+109
View File
@@ -0,0 +1,109 @@
from config.logger import setup_logging
import os
import random
import difflib
import re
import traceback
from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message
TAG = __name__
logger = setup_logging()
def _extract_song_name(text):
"""从用户输入中提取歌名"""
for keyword in ["", "播放", "", ""]:
if keyword in text:
parts = text.split(keyword)
if len(parts) > 1:
return parts[1].strip()
return None
def _find_best_match(potential_song, music_files):
"""查找最匹配的歌曲"""
best_match = None
highest_ratio = 0
for music_file in music_files:
song_name = os.path.splitext(music_file)[0]
ratio = difflib.SequenceMatcher(None, potential_song, song_name).ratio()
if ratio > highest_ratio and ratio > 0.4:
highest_ratio = ratio
best_match = music_file
return best_match
class MusicHandler:
def __init__(self, config):
self.config = config
self.music_related_keywords = []
if "music" in self.config:
self.music_config = self.config["music"]
self.music_dir = os.path.abspath(
self.music_config.get("music_dir", "./music") # 默认路径修改
)
self.music_related_keywords = self.music_config.get("music_commands", [])
else:
self.music_dir = os.path.abspath("./music")
self.music_related_keywords = ["来一首歌", "唱一首歌", "播放音乐", "来点音乐", "背景音乐", "放首歌",
"播放歌曲", "来点背景音乐", "我想听歌", "我要听歌", "放点音乐"]
async def handle_music_command(self, conn, text):
"""处理音乐播放指令"""
clean_text = re.sub(r'[^\w\s]', '', text).strip()
logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}")
# 尝试匹配具体歌名
if os.path.exists(self.music_dir):
music_files = [f for f in os.listdir(self.music_dir) if f.endswith('.mp3')]
logger.bind(tag=TAG).debug(f"找到的音乐文件: {music_files}")
potential_song = _extract_song_name(clean_text)
if potential_song:
best_match = _find_best_match(potential_song, music_files)
if best_match:
logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
await self.play_local_music(conn, specific_file=best_match)
return True
# 检查是否是通用播放音乐命令
if any(cmd in clean_text for cmd in self.music_related_keywords):
await self.play_local_music(conn)
return True
return False
async def play_local_music(self, conn, specific_file=None):
"""播放本地音乐文件"""
try:
if not os.path.exists(self.music_dir):
logger.bind(tag=TAG).error(f"音乐目录不存在: {self.music_dir}")
return
# 确保路径正确性
if specific_file:
music_path = os.path.join(self.music_dir, specific_file)
if not os.path.exists(music_path):
logger.bind(tag=TAG).error(f"指定的音乐文件不存在: {music_path}")
return
selected_music = specific_file
else:
music_files = [f for f in os.listdir(self.music_dir) if f.endswith('.mp3')]
if not music_files:
logger.bind(tag=TAG).error("未找到MP3音乐文件")
return
selected_music = random.choice(music_files)
music_path = os.path.join(self.music_dir, selected_music)
text = f"正在播放{selected_music}"
await send_stt_message(conn, text)
conn.tts_first_text = selected_music
conn.tts_last_text = selected_music
conn.llm_finish_task = True
opus_packets, duration = conn.tts.wav_to_opus_data(music_path)
await sendAudioMessage(conn, opus_packets, duration, selected_music)
except Exception as e:
logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")
logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}")
+81
View File
@@ -0,0 +1,81 @@
from config.logger import setup_logging
import asyncio
import time
from core.utils.util import remove_punctuation_and_length
from core.handle.sendAudioHandle import schedule_with_interrupt, send_stt_message
TAG = __name__
logger = setup_logging()
async def handleAudioMessage(conn, audio):
if not conn.asr_server_receive:
logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
return
if conn.client_listen_mode == "auto":
have_voice = conn.vad.is_vad(conn, audio)
else:
have_voice = conn.client_have_voice
# 如果本次没有声音,本段也没声音,就把声音丢弃了
if have_voice == False and conn.client_have_voice == False:
await no_voice_close_connect(conn)
conn.asr_audio.clear()
return
conn.client_no_voice_last_time = 0.0
conn.asr_audio.append(audio)
# 如果本段有声音,且已经停止了
if conn.client_voice_stop:
conn.client_abort = False
conn.asr_server_receive = False
# 音频太短了,无法识别
if len(conn.asr_audio) < 3:
conn.asr_server_receive = True
else:
text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
logger.bind(tag=TAG).info(f"识别文本: {text}")
text_len, text_without_punctuation = remove_punctuation_and_length(text)
if await conn.music_handler.handle_music_command(conn, text_without_punctuation):
conn.asr_server_receive = True
conn.asr_audio.clear()
return
if text_len <= conn.max_cmd_length and await handleCMDMessage(conn, text_without_punctuation):
return
if text_len > 0:
await startToChat(conn, text)
else:
conn.asr_server_receive = True
conn.asr_audio.clear()
conn.reset_vad_states()
async def handleCMDMessage(conn, text):
cmd_exit = conn.cmd_exit
for cmd in cmd_exit:
if text == cmd:
logger.bind(tag=TAG).info("识别到明确的退出命令".format(text))
await conn.close()
return True
return False
async def startToChat(conn, text):
# 异步发送 stt 信息
stt_task = asyncio.create_task(
schedule_with_interrupt(0, send_stt_message(conn, text))
)
conn.scheduled_tasks.append(stt_task)
conn.executor.submit(conn.chat, text)
async def no_voice_close_connect(conn):
if conn.client_no_voice_last_time == 0.0:
conn.client_no_voice_last_time = time.time() * 1000
else:
no_voice_time = time.time() * 1000 - conn.client_no_voice_last_time
close_connection_no_voice_time = conn.config.get("close_connection_no_voice_time", 120)
if no_voice_time > 1000 * close_connection_no_voice_time:
conn.client_abort = False
conn.asr_server_receive = False
prompt = "时间过得真快,我都好久没说话了。请你用十个字左右话跟我告别,以“再见”或“拜拜”为结尾"
await startToChat(conn, prompt)
@@ -8,56 +8,6 @@ TAG = __name__
logger = setup_logging() logger = setup_logging()
async def handleAudioMessage(conn, audio):
if not conn.asr_server_receive:
logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
return
if conn.client_listen_mode == "auto":
have_voice = conn.vad.is_vad(conn, audio)
else:
have_voice = conn.client_have_voice
# 如果本次没有声音,本段也没声音,就把声音丢弃了
if have_voice == False and conn.client_have_voice == False:
await no_voice_close_connect(conn)
conn.asr_audio.clear()
return
conn.client_no_voice_last_time = 0.0
conn.asr_audio.append(audio)
# 如果本段有声音,且已经停止了
if conn.client_voice_stop:
conn.client_abort = False
conn.asr_server_receive = False
# 音频太短了,无法识别
if len(conn.asr_audio) < 3:
conn.asr_server_receive = True
else:
text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
logger.bind(tag=TAG).info(f"识别文本: {text}")
text_len, text_without_punctuation = remove_punctuation_and_length(text)
if text_len <= conn.max_cmd_length and await handleCMDMessage(conn, text_without_punctuation):
return
if text_len > 0:
await startToChat(conn, text)
else:
conn.asr_server_receive = True
conn.asr_audio.clear()
conn.reset_vad_states()
async def handleCMDMessage(conn, text):
cmd_exit = conn.cmd_exit
for cmd in cmd_exit:
if text == cmd:
logger.bind(tag=TAG).info("识别到明确的退出命令".format(text))
await finishToChat(conn)
return True
return False
async def finishToChat(conn):
await conn.close()
async def isLLMWantToFinish(conn): async def isLLMWantToFinish(conn):
first_text = conn.tts_first_text first_text = conn.tts_first_text
last_text = conn.tts_last_text last_text = conn.tts_last_text
@@ -70,15 +20,6 @@ async def isLLMWantToFinish(conn):
return False return False
async def startToChat(conn, text):
# 异步发送 stt 信息
stt_task = asyncio.create_task(
schedule_with_interrupt(0, send_stt_message(conn, text))
)
conn.scheduled_tasks.append(stt_task)
conn.executor.submit(conn.chat, text)
async def sendAudioMessage(conn, audios, duration, text): async def sendAudioMessage(conn, audios, duration, text):
base_delay = conn.tts_duration base_delay = conn.tts_duration
@@ -107,7 +48,7 @@ async def sendAudioMessage(conn, audios, duration, text):
conn.scheduled_tasks.append(stop_task) conn.scheduled_tasks.append(stop_task)
if await isLLMWantToFinish(conn): if await isLLMWantToFinish(conn):
finish_task = asyncio.create_task( finish_task = asyncio.create_task(
schedule_with_interrupt(stop_duration, finishToChat(conn)) schedule_with_interrupt(stop_duration, await conn.close())
) )
conn.scheduled_tasks.append(finish_task) conn.scheduled_tasks.append(finish_task)
@@ -152,16 +93,3 @@ async def schedule_with_interrupt(delay, coro):
await coro await coro
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
async def no_voice_close_connect(conn):
if conn.client_no_voice_last_time == 0.0:
conn.client_no_voice_last_time = time.time() * 1000
else:
no_voice_time = time.time() * 1000 - conn.client_no_voice_last_time
close_connection_no_voice_time = conn.config.get("close_connection_no_voice_time", 120)
if no_voice_time > 1000 * close_connection_no_voice_time:
conn.client_abort = False
conn.asr_server_receive = False
prompt = "时间过得真快,我都好久没说话了。请你用十个字左右话跟我告别,以“再见”或“拜拜拜”为结尾"
await startToChat(conn, prompt)
+5 -1
View File
@@ -2,7 +2,8 @@ from config.logger import setup_logging
import json import json
from core.handle.abortHandle import handleAbortMessage from core.handle.abortHandle import handleAbortMessage
from core.handle.helloHandle import handleHelloMessage from core.handle.helloHandle import handleHelloMessage
from core.handle.audioHandle import startToChat from core.handle.receiveAudioHandle import startToChat
from core.handle.iotHandle import handleIotDescriptors
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
@@ -36,5 +37,8 @@ async def handleTextMessage(conn, message):
conn.asr_audio.clear() conn.asr_audio.clear()
if "text" in msg_json: if "text" in msg_json:
await startToChat(conn, msg_json["text"]) await startToChat(conn, msg_json["text"])
elif msg_json["type"] == "iot":
if "descriptors" in msg_json:
await handleIotDescriptors(conn, msg_json["descriptors"])
except json.JSONDecodeError: except json.JSONDecodeError:
await conn.websocket.send(message) await conn.websocket.send(message)
+5 -5
View File
@@ -8,7 +8,7 @@ import websockets
import json import json
import gzip import gzip
import opuslib import opuslib_next
from core.providers.asr.base import ASRProviderBase from core.providers.asr.base import ASRProviderBase
from config.logger import setup_logging from config.logger import setup_logging
@@ -103,14 +103,14 @@ class ASRProvider(ASRProviderBase):
file_name = f"asr_{session_id}_{uuid.uuid4()}.wav" file_name = f"asr_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
decoder = opuslib.Decoder(16000, 1) # 16kHz, 单声道 decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = [] pcm_data = []
for opus_packet in opus_data: for opus_packet in opus_data:
try: try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame) pcm_data.append(pcm_frame)
except opuslib.OpusError as e: except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True) logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
@@ -216,14 +216,14 @@ class ASRProvider(ASRProviderBase):
@staticmethod @staticmethod
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]: def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
decoder = opuslib.Decoder(16000, 1) # 16kHz, 单声道 decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = [] pcm_data = []
for opus_packet in opus_data: for opus_packet in opus_data:
try: try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame) pcm_data.append(pcm_frame)
except opuslib.OpusError as e: except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True) logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data return pcm_data
+18 -5
View File
@@ -25,12 +25,25 @@ class LLMProvider(LLMProviderBase):
messages=dialogue, messages=dialogue,
stream=True stream=True
) )
is_active = True
for chunk in responses: for chunk in responses:
# 检查是否存在有效的choice且content不为空 try:
if chunk.choices and len(chunk.choices) > 0: # 检查是否存在有效的choice且content不为空
delta = chunk.choices[0].delta delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None
content = getattr(delta, 'content', '') content = delta.content if hasattr(delta, 'content') else ''
if content: # 仅在content非空时生成 except IndexError:
content = ''
if content:
# 处理标签跨多个chunk的情况
if '<think>' in content:
is_active = False
content = content.split('<think>')[0]
if '</think>' in content:
is_active = True
content = content.split('</think>')[-1]
if is_active:
yield content yield content
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error in response generation: {e}") logger.bind(tag=TAG).error(f"Error in response generation: {e}")
+112 -28
View File
@@ -1,5 +1,13 @@
import os import os
import sys import sys
import asyncio
from typing import List, Dict, Any
# 添加项目根目录到Python路径
current_dir = os.path.dirname(os.path.abspath(__file__))
project_root = os.path.abspath(os.path.join(current_dir, "..", ".."))
sys.path.insert(0, project_root)
from config.logger import setup_logging from config.logger import setup_logging
import importlib import importlib
from datetime import datetime from datetime import datetime
@@ -21,39 +29,115 @@ def create_instance(class_name, *args, **kwargs):
raise ValueError(f"不支持的LLM类型: {class_name},请检查该配置的type是否设置正确") raise ValueError(f"不支持的LLM类型: {class_name},请检查该配置的type是否设置正确")
if __name__ == "__main__": async def test_single_model(llm_name: str, llm_config: Dict[str, Any], test_prompt: str, config: Dict[str, Any]) -> Dict[str, Any]:
"""异步测试单个模型"""
try:
# 获取实际的LLM类型
llm_type = llm_config["type"] if "type" in llm_config else llm_name
llm = create_instance(llm_type, llm_config)
# 开始测试
dialogue = []
dialogue.append({"role": "system", "content": config.get("prompt")})
dialogue.append({"role": "user", "content": test_prompt})
start_time = datetime.now()
llm_responses = llm.response("test", dialogue)
response_message = []
first_response_time = None
total_response_time = None
start = 0
full_response = ""
for content in llm_responses:
response_message.append(content)
full_response += content
if is_segment(response_message):
segment_text = "".join(response_message[start:])
segment_text = get_string_no_punctuation_or_emoji(segment_text)
if len(segment_text) > 0:
if first_response_time is None:
first_response_time = (datetime.now() - start_time).total_seconds()
start = len(response_message)
total_response_time = (datetime.now() - start_time).total_seconds()
return {
"name": llm_name,
"type": llm_type,
"first_response_time": first_response_time,
"total_response_time": total_response_time,
"response_length": len(full_response),
"status": "成功",
"response": full_response
}
except Exception as e:
print(f"测试 {llm_name} 时发生错误: {str(e)}")
return {
"name": llm_name,
"type": llm_config.get("type", llm_name),
"first_response_time": None,
"total_response_time": None,
"response_length": 0,
"status": f"失败 - {str(e)}",
"response": ""
}
async def main():
""" """
响应速度测试 LLM模型响应速度测试和排行(异步版本)
""" """
config = read_config(get_project_dir() + "config.yaml") config = read_config(get_project_dir() + "config.yaml")
llm = create_instance( test_prompt = "你好小智"
config["selected_module"]["LLM"]
if not "type" in config["LLM"][config["selected_module"]["LLM"]] print("开始并发测试所有模型...")
else
config["LLM"][config["selected_module"]["LLM"]]["type"], # 创建所有模型的测试任务
config["LLM"][config["selected_module"]["LLM"]] tasks = []
) for llm_name, llm_config in config["LLM"].items():
task = asyncio.create_task(test_single_model(llm_name, llm_config, test_prompt, config))
tasks.append(task)
# 等待所有测试完成
test_results = await asyncio.gather(*tasks)
# 打印测试结果排行榜
print("\n========= LLM模型性能测试排行榜 =========")
print("测试提示词:", test_prompt)
# 过滤出成功的结果,并确保数值有效
successful_results = [r for r in test_results if r["status"] == "成功" and r["first_response_time"] is not None]
if successful_results:
print("\n1. 首次响应时间排行:")
sorted_by_first = sorted(successful_results, key=lambda x: x["first_response_time"])
for i, result in enumerate(sorted_by_first, 1):
print(f"{i}. {result['name']}({result['type']}) - {result['first_response_time']:.2f}")
print(f" 响应内容: {result['response'][:50]}...") # 只显示前50个字符
start_time = datetime.now() print("\n2. 总响应时间排行:")
sorted_by_total = sorted(successful_results, key=lambda x: x["total_response_time"] or float('inf'))
for i, result in enumerate(sorted_by_total, 1):
if result["total_response_time"] is not None:
print(f"{i}. {result['name']}({result['type']}) - {result['total_response_time']:.2f}")
dialogue = [] print("\n3. 响应长度比较:")
dialogue.append({"role": "system", "content": config.get("prompt")}) sorted_by_length = sorted(successful_results, key=lambda x: x["response_length"], reverse=True)
dialogue.append({"role": "user", "content": "你好小智"}) for i, result in enumerate(sorted_by_length, 1):
llm_responses = llm.response("test", dialogue) print(f"{i}. {result['name']}({result['type']}) - {result['response_length']}字符")
response_message = [] else:
first_text = None print("\n没有成功完成测试的模型。")
start = 0
for content in llm_responses: if len(test_results) != len(successful_results):
response_message.append(content) print("\n测试失败的模型:")
failed_results = [r for r in test_results if r["status"] != "成功" or r["first_response_time"] is None]
for result in failed_results:
print(f"- {result['name']}({result['type']}): {result['status']}")
if is_segment(response_message):
segment_text = "".join(response_message[start:])
segment_text = get_string_no_punctuation_or_emoji(segment_text)
if len(segment_text) > 0:
if first_text is None:
first_text = segment_text
print("大模型首次返回耗时:" + str(datetime.now() - start_time))
start = len(response_message)
print("大模型返回总耗时:" + str(datetime.now() - start_time)) if __name__ == "__main__":
# 运行异步主函数
asyncio.run(main())
+2 -2
View File
@@ -90,7 +90,7 @@ def remove_punctuation_and_length(text):
char not in full_width_punctuations and char not in half_width_punctuations and char not in space and char not in full_width_space]) char not in full_width_punctuations and char not in half_width_punctuations and char not in space and char not in full_width_space])
if result == "Yeah": if result == "Yeah":
return 0 return 0, ""
return len(result), result return len(result), result
@@ -120,4 +120,4 @@ def check_password(password):
return False return False
# 如果满足所有条件,则返回True # 如果满足所有条件,则返回True
return True return True
+12 -4
View File
@@ -2,6 +2,7 @@ import asyncio
import websockets import websockets
from config.logger import setup_logging from config.logger import setup_logging
from core.connection import ConnectionHandler from core.connection import ConnectionHandler
from core.handle.musicHandler import MusicHandler
from core.utils.util import get_local_ip from core.utils.util import get_local_ip
from core.utils import asr, vad, llm, tts from core.utils import asr, vad, llm, tts
@@ -12,7 +13,8 @@ class WebSocketServer:
def __init__(self, config: dict): def __init__(self, config: dict):
self.config = config self.config = config
self.logger = setup_logging() self.logger = setup_logging()
self._vad, self._asr, self._llm, self._tts = self._create_processing_instances() self._vad, self._asr, self._llm, self._tts, self._music = self._create_processing_instances()
self.active_connections = set() # 添加全局连接记录
def _create_processing_instances(self): def _create_processing_instances(self):
"""创建处理模块实例""" """创建处理模块实例"""
@@ -43,7 +45,8 @@ class WebSocketServer:
self.config["TTS"][self.config["selected_module"]["TTS"]]["type"], 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"]
) ),
MusicHandler(self.config)
) )
async def start(self): async def start(self):
@@ -62,5 +65,10 @@ class WebSocketServer:
async def _handle_connection(self, websocket): async def _handle_connection(self, websocket):
"""处理新连接,每次创建独立的ConnectionHandler""" """处理新连接,每次创建独立的ConnectionHandler"""
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts) # 创建ConnectionHandler时传入当前server实例
await handler.handle_connection(websocket) handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music)
self.active_connections.add(handler)
try:
await handler.handle_connection(websocket)
finally:
self.active_connections.discard(handler)
+1
View File
@@ -1,3 +1,4 @@
version: '3'
services: services:
xiaozhi-esp32-server: xiaozhi-esp32-server:
image: ghcr.nju.edu.cn/xinnan-tech/xiaozhi-esp32-server:latest image: ghcr.nju.edu.cn/xinnan-tech/xiaozhi-esp32-server:latest
+3 -3
View File
@@ -117,7 +117,7 @@ docker run -it --name xiaozhi-env --restart always --security-opt seccomp:unconf
-p 8000:8000 \ -p 8000:8000 \
-p 8002:8002 \ -p 8002:8002 \
-v ./:/app \ -v ./:/app \
ccr.ccs.tencentyun.com/kalicyh/poetry:v3.10_latest kalicyh/poetry:v3.10_xiaozhi
``` ```
然后就和正常开发一样了 然后就和正常开发一样了
@@ -163,8 +163,8 @@ poetry run python app.py
conda remove -n xiaozhi-esp32-server --all -y conda remove -n xiaozhi-esp32-server --all -y
conda create -n xiaozhi-esp32-server python=3.10 -y conda create -n xiaozhi-esp32-server python=3.10 -y
conda activate xiaozhi-esp32-server conda activate xiaozhi-esp32-server
conda install conda-forge::libopus conda install conda-forge::libopus -y
conda install conda-forge::ffmpeg conda install conda-forge::ffmpeg -y
``` ```
## 2.安装本项目依赖 ## 2.安装本项目依赖
Binary file not shown.

Before

Width:  |  Height:  |  Size: 446 KiB

After

Width:  |  Height:  |  Size: 427 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.
+328 -179
View File
@@ -1,27 +1,27 @@
import time import time
import aiohttp
import asyncio
from tabulate import tabulate from tabulate import tabulate
from typing import Dict from typing import Dict, List
from core.utils.llm import create_instance as create_llm_instance from core.utils.llm import create_instance as create_llm_instance
from core.utils.tts import create_instance as create_tts_instance from core.utils.tts import create_instance as create_tts_instance
from core.utils.util import read_config from core.utils.util import read_config
import statistics import statistics
from config.settings import get_config_file from config.settings import get_config_file
from concurrent.futures import ThreadPoolExecutor
import inspect import inspect
import os import os
import requests
import logging import logging
# 设置全局日志级别为WARNING,抑制INFO级别日志 # 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
class PerformanceTester:
class AsyncPerformanceTester:
def __init__(self): def __init__(self):
self.config = read_config(get_config_file()) self.config = read_config(get_config_file())
# 从配置读取测试句子,如果不存在则使用默认
self.test_sentences = self.config.get("module_test", {}).get( self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences", "test_sentences",
["你好,请介绍一下你自己", "What's the weather like today?", ["你好,请介绍一下你自己", "What's the weather like today?",
"请用100字概括量子计算的基本原理和应用前景"] "请用100字概括量子计算的基本原理和应用前景"]
) )
self.results = { self.results = {
@@ -30,258 +30,407 @@ class PerformanceTester:
"combinations": [] "combinations": []
} }
def _test_llm(self, llm_name: str, config: Dict) -> Dict: async def _check_ollama_service(self, base_url: str, model_name: str) -> bool:
"""测试单个LLM性能""" """异步检查Ollama服务状态"""
async with aiohttp.ClientSession() as session:
try:
# 检查服务是否可用
async with session.get(f"{base_url}/api/version") as response:
if response.status != 200:
print(f"🚫 Ollama服务未启动或无法访问: {base_url}")
return False
# 检查模型是否存在
async with session.get(f"{base_url}/api/tags") as response:
if response.status == 200:
data = await response.json()
models = data.get("models", [])
if not any(model["name"] == model_name for model in models):
print(f"🚫 Ollama模型 {model_name} 未找到,请先使用 ollama pull {model_name} 下载")
return False
else:
print(f"🚫 无法获取Ollama模型列表")
return False
return True
except Exception as e:
print(f"🚫 无法连接到Ollama服务: {str(e)}")
return False
async def _test_tts(self, tts_name: str, config: Dict) -> Dict:
"""异步测试单个TTS性能"""
try: try:
# 跳过未配置密钥的模块 logging.getLogger("core.providers.tts.base").setLevel(logging.WARNING)
if "api_key" in config and any(x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"]):
print(f"🚫 跳过未配置的LLM: {llm_name}") token_fields = ["access_token", "api_key", "token"]
return {"errors": 1} if any(field in config and any(x in config[field] for x in ["你的", "placeholder"]) for field in
token_fields):
print(f"⏭️ TTS {tts_name} 未配置access_token/api_key,已跳过")
return {"name": tts_name, "type": "tts", "errors": 1}
module_type = config.get('type', tts_name)
tts = create_tts_instance(
module_type,
config,
delete_audio_file=True
)
print(f"🎵 测试 TTS: {tts_name}")
tmp_file = tts.generate_filename()
await tts.text_to_speak("连接测试", tmp_file)
if not tmp_file or not os.path.exists(tmp_file):
print(f"{tts_name} 连接失败")
return {"name": tts_name, "type": "tts", "errors": 1}
total_time = 0
test_count = len(self.test_sentences[:2])
for i, sentence in enumerate(self.test_sentences[:2], 1):
start = time.time()
tmp_file = tts.generate_filename()
await tts.text_to_speak(sentence, tmp_file)
duration = time.time() - start
total_time += duration
if tmp_file and os.path.exists(tmp_file):
print(f"{tts_name} [{i}/{test_count}]")
else:
print(f"{tts_name} [{i}/{test_count}]")
return {"name": tts_name, "type": "tts", "errors": 1}
return {
"name": tts_name,
"type": "tts",
"avg_time": total_time / test_count,
"errors": 0
}
except Exception as e:
print(f"⚠️ {tts_name} 测试失败: {str(e)}")
return {"name": tts_name, "type": "tts", "errors": 1}
async def _test_llm(self, llm_name: str, config: Dict) -> Dict:
"""异步测试单个LLM性能"""
try:
# 对于Ollama,跳过api_key检查并进行特殊处理
if llm_name == "Ollama":
base_url = config.get('base_url', 'http://localhost:11434')
model_name = config.get('model_name')
if not model_name:
print(f"🚫 Ollama未配置model_name")
return {"name": llm_name, "type": "llm", "errors": 1}
if not await self._check_ollama_service(base_url, model_name):
return {"name": llm_name, "type": "llm", "errors": 1}
else:
if "api_key" in config and any(x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"]):
print(f"🚫 跳过未配置的LLM: {llm_name}")
return {"name": llm_name, "type": "llm", "errors": 1}
# 获取实际类型(兼容旧配置) # 获取实际类型(兼容旧配置)
module_type = config.get('type', llm_name) module_type = config.get('type', llm_name)
llm = create_llm_instance(module_type, config) llm = create_llm_instance(module_type, config)
# 统一使用UTF-8编码 # 统一使用UTF-8编码
test_sentences = [s.encode('utf-8').decode('utf-8') for s in self.test_sentences] test_sentences = [s.encode('utf-8').decode('utf-8') for s in self.test_sentences]
total_time = 0 # 创建所有句子的测试任务
first_token_times = [] sentence_tasks = []
valid_times = []
for sentence in test_sentences: for sentence in test_sentences:
sentence_start = time.time() # 记录整句开始时间 sentence_tasks.append(self._test_single_sentence(llm_name, llm, sentence))
first_token_received = False
# 并发执行所有句子测试
# 遍历响应流 sentence_results = await asyncio.gather(*sentence_tasks)
for chunk in llm.response("perf_test", [{"role": "user", "content": sentence}]):
if not first_token_received and chunk.strip() != '': # 处理结果
first_token_times.append(time.time() - sentence_start) valid_results = [r for r in sentence_results if r is not None]
first_token_received = True if not valid_results:
# 计算整句耗时
sentence_duration = time.time() - sentence_start
total_time += sentence_duration
valid_times.append(sentence_duration)
# 新增有效性检查
if len(first_token_times) == 0 or len(valid_times) == 0:
print(f"⚠️ {llm_name} 无有效数据,可能配置错误") print(f"⚠️ {llm_name} 无有效数据,可能配置错误")
return {"errors": 1} return {"name": llm_name, "type": "llm", "errors": 1}
# 过滤异常数据(超过3倍标准差) first_token_times = [r["first_token_time"] for r in valid_results]
mean = statistics.mean(valid_times) response_times = [r["response_time"] for r in valid_results]
stdev = statistics.stdev(valid_times) if len(valid_times) > 1 else 0
filtered_times = [t for t in valid_times if t <= mean + 3*stdev] # 过滤异常数据
mean = statistics.mean(response_times)
# 当有效数据不足时标记错误 stdev = statistics.stdev(response_times) if len(response_times) > 1 else 0
filtered_times = [t for t in response_times if t <= mean + 3 * stdev]
if len(filtered_times) < len(test_sentences) * 0.5: if len(filtered_times) < len(test_sentences) * 0.5:
print(f"⚠️ {llm_name} 有效数据不足,可能网络不稳定") print(f"⚠️ {llm_name} 有效数据不足,可能网络不稳定")
return {"errors": 1} return {"name": llm_name, "type": "llm", "errors": 1}
return { return {
"avg_response": total_time / len(test_sentences), "name": llm_name,
"avg_first_token": sum(first_token_times)/len(first_token_times), "type": "llm",
"avg_response": sum(response_times) / len(response_times),
"avg_first_token": sum(first_token_times) / len(first_token_times),
"std_first_token": statistics.stdev(first_token_times) if len(first_token_times) > 1 else 0, "std_first_token": statistics.stdev(first_token_times) if len(first_token_times) > 1 else 0,
"std_response": statistics.stdev(valid_times) if len(valid_times) > 1 else 0, "std_response": statistics.stdev(response_times) if len(response_times) > 1 else 0,
"errors": 0 "errors": 0
} }
except Exception as e: except Exception as e:
print(f"LLM {llm_name} 测试失败: {str(e)}") print(f"LLM {llm_name} 测试失败: {str(e)}")
return {"errors": 1} return {"name": llm_name, "type": "llm", "errors": 1}
def _test_tts(self, tts_name: str, config: Dict) -> Dict: async def _test_single_sentence(self, llm_name: str, llm, sentence: str) -> Dict:
"""测试单个TTS性能""" """测试单个句子的性能"""
try: try:
# 关闭详细日志 print(f"📝 {llm_name} 开始测试: {sentence[:20]}...")
logging.getLogger("core.providers.tts.base").setLevel(logging.WARNING) sentence_start = time.time()
first_token_received = False
# 跳过未配置密钥的模块 first_token_time = None
token_fields = ["access_token", "api_key", "token"]
if any(field in config and any(x in config[field] for x in ["你的", "placeholder"]) for field in token_fields):
print(f"⏭️ TTS {tts_name} 未配置access_token/api_key,已跳过")
return {"errors": 1}
# 获取实际类型(兼容旧配置)
module_type = config.get('type', tts_name)
tts = create_tts_instance(
module_type,
config,
delete_audio_file=True # 确保参数名称正确
)
# 简化后的输出
print(f"\n🎵 正在测试 TTS: {tts_name}")
print(f"🔊 测试 {tts_name}", end="", flush=True)
# 连接测试
test_conn = tts.to_tts("连接测试")
if not os.path.exists(test_conn):
print("❌ 连接失败")
return {"errors": 1}
else:
print("")
total_time = 0
test_count = len(self.test_sentences[:2])
for i, sentence in enumerate(self.test_sentences[:2], 1):
start = time.time()
file_path = tts.to_tts(sentence)
duration = time.time() - start
total_time += duration
# 显示简单的进度标识
if os.path.exists(file_path):
print(f"✓[{i}/{test_count}]", end="", flush=True)
else:
print(f"✗[{i}/{test_count}]", end="", flush=True)
print() # 换行
return {"avg_time": total_time / test_count, "errors": 0}
except requests.exceptions.ConnectionError:
print(f"\n{tts_name} 无法连接服务端")
return {"errors": 1}
except Exception as e:
print(f"\n⚠️ {tts_name} 测试失败: {str(e)}")
return {"errors": 1}
def run(self): async def process_response():
"""执行全量测试并自动跳过未配置的模块""" nonlocal first_token_received, first_token_time
print("🔍 开始自动检测已配置的模块...") for chunk in llm.response("perf_test", [{"role": "user", "content": sentence}]):
if not first_token_received and chunk.strip() != '':
# 测试所有LLM first_token_time = time.time() - sentence_start
for llm_name, config in self.config.get("LLM", {}).items(): first_token_received = True
# 特殊处理CozeLLM的配置检查 print(f"{llm_name} 首个Token: {first_token_time:.3f}s")
if llm_name == "CozeLLM": yield chunk
if any(x in config.get("bot_id", "") for x in ["你的"]) \
or any(x in config.get("user_id", "") for x in ["你的"]): response_chunks = []
print(f"⏭️ LLM {llm_name} 未配置bot_id/user_id,已跳过") async for chunk in process_response():
continue response_chunks.append(chunk)
# 通用的api_key配置检查
if "api_key" in config and any(x in config["api_key"] for x in ["你的", "placeholder"]): response_time = time.time() - sentence_start
print(f"⏭️ LLM {llm_name} 未配置api_key,已跳过") print(f" {llm_name} 完成响应: {response_time:.3f}s")
continue
if first_token_time is None:
print(f"🚀 正在测试 LLM: {llm_name}") first_token_time = response_time # 如果没有检测到first token,使用总响应时间
self.results["llm"][llm_name] = self._test_llm(llm_name, config)
return {
# 测试所有TTS "name": llm_name,
for tts_name, config in self.config.get("TTS", {}).items(): "type": "llm",
# 根据不同服务的token字段检测 "first_token_time": first_token_time,
token_fields = ["access_token", "api_key", "token"] "response_time": response_time
if any(field in config and any(x in config[field] for x in ["你的", "placeholder"]) for field in token_fields): }
print(f"⏭️ TTS {tts_name} 未配置access_token/api_key,已跳过") except Exception as e:
continue print(f"⚠️ {llm_name} 句子测试失败: {str(e)}")
return None
print(f"🎵 正在测试 TTS: {tts_name}")
self.results["tts"][tts_name] = self._test_tts(tts_name, config)
# 生成组合建议
self._generate_combinations()
self._print_results()
def _generate_combinations(self): def _generate_combinations(self):
"""生成最佳组合建议""" """生成最佳组合建议"""
# 调整过滤条件,例如设为 >= 0.05
valid_llms = [ valid_llms = [
k for k, v in self.results["llm"].items() k for k, v in self.results["llm"].items()
if v["errors"] == 0 and v["avg_first_token"] >= 0.05 if v["errors"] == 0 and v["avg_first_token"] >= 0.05
] ]
valid_tts = [k for k, v in self.results["tts"].items() if v["errors"] == 0] valid_tts = [k for k, v in self.results["tts"].items() if v["errors"] == 0]
# 找出基准值
min_first_token = min([self.results["llm"][llm]["avg_first_token"] for llm in valid_llms]) if valid_llms else 1
min_tts_time = min([self.results["tts"][tts]["avg_time"] for tts in valid_tts]) if valid_tts else 1
for llm in valid_llms: for llm in valid_llms:
for tts in valid_tts: for tts in valid_tts:
llm_weight = 0.8 if self.results["llm"][llm]["avg_first_token"] < 1.0 else 0.6 # 计算相对性能分数(越小越好)
tts_weight = 1 - llm_weight llm_score = self.results["llm"][llm]["avg_first_token"] / min_first_token
score = ( tts_score = self.results["tts"][tts]["avg_time"] / min_tts_time
self.results["llm"][llm]["avg_first_token"] * llm_weight +
self.results["tts"][tts]["avg_time"] * tts_weight # 计算稳定性分数(标准差/平均值,越小越稳定)
) llm_stability = self.results["llm"][llm]["std_first_token"] / self.results["llm"][llm][
"avg_first_token"]
# 综合得分(考虑性能和稳定性)
# 性能权重0.7,稳定性权重0.3
llm_final_score = llm_score * 0.7 + llm_stability * 0.3
# 总分 = LLM得分(70%) + TTS得分(30%)
total_score = llm_final_score * 0.7 + tts_score * 0.3
self.results["combinations"].append({ self.results["combinations"].append({
"llm": llm, "llm": llm,
"tts": tts, "tts": tts,
"score": score, "score": total_score,
"details": { "details": {
"llm_first_token": self.results["llm"][llm]["avg_first_token"], "llm_first_token": self.results["llm"][llm]["avg_first_token"],
"llm_stability": llm_stability,
"tts_time": self.results["tts"][tts]["avg_time"] "tts_time": self.results["tts"][tts]["avg_time"]
} }
}) })
# 按综合得分排序 # 分数越小越好
self.results["combinations"].sort(key=lambda x: x["score"]) self.results["combinations"].sort(key=lambda x: x["score"])
def _print_results(self): def _print_results(self):
"""控制台输出结果""" """打印测试结果"""
# LLM结果表格
llm_table = [] llm_table = []
for name, data in self.results["llm"].items(): for name, data in self.results["llm"].items():
if data["errors"] == 0: if data["errors"] == 0:
stability = data["std_first_token"] / data["avg_first_token"]
llm_table.append([ llm_table.append([
name, name, # 不需要固定宽度,让tabulate自己处理对齐
f"{data['avg_first_token']:.3f}s", f"{data['avg_first_token']:.3f}",
f"{data['avg_response']:.3f}s" f"{data['avg_response']:.3f}",
f"{stability:.3f}"
]) ])
if llm_table: if llm_table:
print("\nLLM 性能排行:") print("\nLLM 性能排行:")
print(tabulate( print(tabulate(
llm_table, llm_table,
headers=["名称", "平均首Token时间", "平均总响应时间"], headers=["名称", "首字耗时", "总耗时", "稳定性"],
tablefmt="github" tablefmt="github",
colalign=("left", "right", "right", "right"),
disable_numparse=True
)) ))
else: else:
print("\n⚠️ 没有可用的LLM模块进行测试。") print("\n⚠️ 没有可用的LLM模块进行测试。")
# TTS结果表格
tts_table = [] tts_table = []
for name, data in self.results["tts"].items(): for name, data in self.results["tts"].items():
if data["errors"] == 0: if data["errors"] == 0:
tts_table.append([ tts_table.append([
name, name, # 不需要固定宽度
f"{data['avg_time']:.3f}s" f"{data['avg_time']:.3f}"
]) ])
if tts_table: if tts_table:
print("\nTTS 性能排行:") print("\nTTS 性能排行:")
print(tabulate( print(tabulate(
tts_table, tts_table,
headers=["名称", "平均合成时间"], headers=["名称", "合成耗时"],
tablefmt="github" tablefmt="github",
colalign=("left", "right"),
disable_numparse=True
)) ))
else: else:
print("\n⚠️ 没有可用的TTS模块进行测试。") print("\n⚠️ 没有可用的TTS模块进行测试。")
# 最佳组合建议
if self.results["combinations"]: if self.results["combinations"]:
print("\n推荐配置组合 (综合响应速度):") print("\n推荐配置组合 (得分越小越好):")
combo_table = [] combo_table = []
for combo in self.results["combinations"][:5]: # 显示前5名 for combo in self.results["combinations"][:5]:
combo_table.append([ combo_table.append([
f"{combo['llm']} + {combo['tts']}", f"{combo['llm']} + {combo['tts']}", # 不需要固定宽度
f"{combo['score']:.3f}", f"{combo['score']:.3f}",
f"{combo['details']['llm_first_token']:.3f}s", f"{combo['details']['llm_first_token']:.3f}",
f"{combo['details']['tts_time']:.3f}s" f"{combo['details']['llm_stability']:.3f}",
f"{combo['details']['tts_time']:.3f}"
]) ])
print(tabulate( print(tabulate(
combo_table, combo_table,
headers=["组合方案", "综合得分", "LLM首Token", "TTS合成"], headers=["组合方案", "综合得分", "LLM首字耗时", "稳定性", "TTS合成耗时"],
tablefmt="github" tablefmt="github",
colalign=("left", "right", "right", "right", "right"),
disable_numparse=True
)) ))
else: else:
print("\n⚠️ 没有可用的模块组合建议。") print("\n⚠️ 没有可用的模块组合建议。")
def _execute_with_timeout(self, func, args=(), kwargs={}, timeout=None): def _process_results(self, all_results):
with ThreadPoolExecutor(max_workers=1) as executor: """处理测试结果"""
future = executor.submit(func, *args, **kwargs) for result in all_results:
try: if result["errors"] == 0:
result = future.result(timeout) if result["type"] == "llm":
return list(result) if inspect.isgenerator(result) else result self.results["llm"][result["name"]] = result
except TimeoutError: else:
raise Exception("操作超时") self.results["tts"][result["name"]] = result
async def run(self):
"""执行全量异步测试"""
print("🔍 开始筛选可用模块...")
# 创建所有测试任务
all_tasks = []
# LLM测试任务
for llm_name, config in self.config.get("LLM", {}).items():
# 检查配置有效性
if llm_name == "CozeLLM":
if any(x in config.get("bot_id", "") for x in ["你的"]) \
or any(x in config.get("user_id", "") for x in ["你的"]):
print(f"⏭️ LLM {llm_name} 未配置bot_id/user_id,已跳过")
continue
elif "api_key" in config and any(x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"]):
print(f"⏭️ LLM {llm_name} 未配置api_key,已跳过")
continue
# 对于Ollama,先检查服务状态
if llm_name == "Ollama":
base_url = config.get('base_url', 'http://localhost:11434')
model_name = config.get('model_name')
if not model_name:
print(f"🚫 Ollama未配置model_name")
continue
if not await self._check_ollama_service(base_url, model_name):
continue
print(f"📋 添加LLM测试任务: {llm_name}")
module_type = config.get('type', llm_name)
llm = create_llm_instance(module_type, config)
# 为每个句子创建独立任务
for sentence in self.test_sentences:
sentence = sentence.encode('utf-8').decode('utf-8')
all_tasks.append(self._test_single_sentence(llm_name, llm, sentence))
# TTS测试任务
for tts_name, config in self.config.get("TTS", {}).items():
token_fields = ["access_token", "api_key", "token"]
if any(field in config and any(x in config[field] for x in ["你的", "placeholder"]) for field in
token_fields):
print(f"⏭️ TTS {tts_name} 未配置access_token/api_key,已跳过")
continue
print(f"🎵 添加TTS测试任务: {tts_name}")
all_tasks.append(self._test_tts(tts_name, config))
print(
f"\n✅ 找到 {len([t for t in all_tasks if 'test_single_sentence' in str(t)]) / len(self.test_sentences):.0f} 个可用LLM模块")
print(f"✅ 找到 {len([t for t in all_tasks if '_test_tts' in str(t)])} 个可用TTS模块")
print("\n⏳ 开始并发测试所有模块...\n")
# 并发执行所有测试任务
all_results = await asyncio.gather(*all_tasks, return_exceptions=True)
# 处理LLM结果
llm_results = {}
for result in [r for r in all_results if r and isinstance(r, dict) and r.get("type") == "llm"]:
llm_name = result["name"]
if llm_name not in llm_results:
llm_results[llm_name] = {
"name": llm_name,
"type": "llm",
"first_token_times": [],
"response_times": [],
"errors": 0
}
llm_results[llm_name]["first_token_times"].append(result["first_token_time"])
llm_results[llm_name]["response_times"].append(result["response_time"])
# 计算LLM平均值和标准差
for llm_name, data in llm_results.items():
if len(data["first_token_times"]) >= len(self.test_sentences) * 0.5:
self.results["llm"][llm_name] = {
"name": llm_name,
"type": "llm",
"avg_response": sum(data["response_times"]) / len(data["response_times"]),
"avg_first_token": sum(data["first_token_times"]) / len(data["first_token_times"]),
"std_first_token": statistics.stdev(data["first_token_times"]) if len(
data["first_token_times"]) > 1 else 0,
"std_response": statistics.stdev(data["response_times"]) if len(data["response_times"]) > 1 else 0,
"errors": 0
}
# 处理TTS结果
for result in [r for r in all_results if r and isinstance(r, dict) and r.get("type") == "tts"]:
if result["errors"] == 0:
self.results["tts"][result["name"]] = result
# 生成组合建议并打印结果
print("\n📊 生成测试报告...")
self._generate_combinations()
self._print_results()
async def main():
tester = AsyncPerformanceTester()
await tester.run()
if __name__ == "__main__": if __name__ == "__main__":
tester = PerformanceTester() asyncio.run(main())
tester.run()
+1 -1
View File
@@ -16,4 +16,4 @@ aiohttp_cors==0.7.0
ormsgpack==1.7.0 ormsgpack==1.7.0
ruamel.yaml==0.18.10 ruamel.yaml==0.18.10
loguru==0.7.3 loguru==0.7.3
requests>=2.0.0 requests==2.32.3