Merge pull request #2048 from xinnan-tech/hot-fix

update:优化测试工具目录结构
This commit is contained in:
欣南科技
2025-08-16 10:10:17 +08:00
committed by GitHub
8 changed files with 58 additions and 753 deletions
+1 -2
View File
@@ -227,8 +227,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| 工具名称 | 位置 | 使用方法 | 功能说明 | | 工具名称 | 位置 | 使用方法 | 功能说明 |
|:---:|:---|:---:|:---:| |:---:|:---|:---:|:---:|
| 音频交互测试工具 | main》xiaozhi-server》test》test_page.html | 使用谷歌浏览器直接打开 | 测试音频播放和接收功能,验证Python端音频处理是否正常 | | 音频交互测试工具 | main》xiaozhi-server》test》test_page.html | 使用谷歌浏览器直接打开 | 测试音频播放和接收功能,验证Python端音频处理是否正常 |
| 模型响应测试工具1 | main》xiaozhi-server》performance_tester.py | 执行 `python performance_tester.py` | 测试ASR(语音识别)、LLM(大模型)、TTS(语音合成)三个核心模块的响应速度 | | 模型响应测试工具 | main》xiaozhi-server》performance_tester.py | 执行 `python performance_tester.py` | 测试ASR(语音识别)、LLM(大模型)、VLLM(视觉模型)、TTS(语音合成)三个核心模块的响应速度 |
| 模型响应测试工具2 | main》xiaozhi-server》performance_tester_vllm.py | 执行 `python performance_tester_vllm.py` | 测试VLLM(视觉模型)的响应速度 |
> 💡 提示:测试模型速度时,只会测试配置了密钥的模型。 > 💡 提示:测试模型速度时,只会测试配置了密钥的模型。
+1 -1
View File
@@ -21,7 +21,7 @@ VLLM:
ASR: ASR:
``` ```
4.在main/xiaozhi-server目录下运行performance_text_tool.py: 4.在main/xiaozhi-server目录下运行performance_tester.py:
``` ```
python performance_test_tool.py python performance_test_tool.py
``` ```
@@ -2,17 +2,24 @@ import os
import importlib.util import importlib.util
import asyncio import asyncio
print("使用前请根据doc/performance_texter.md的说明准备配置。") print("使用前请根据doc/performance_testerer.md的说明准备配置。")
def list_performance_text_modules():
performance_text_dir = os.path.join(os.path.dirname(__file__), "performance_text")
def list_performance_tester_modules():
performance_tester_dir = os.path.join(
os.path.dirname(__file__), "performance_tester"
)
modules = [] modules = []
for file in os.listdir(performance_text_dir): for file in os.listdir(performance_tester_dir):
if file.endswith(".py"): if file.endswith(".py"):
modules.append(file[:-3]) modules.append(file[:-3])
return modules return modules
async def load_and_execute_module(module_name): async def load_and_execute_module(module_name):
module_path = os.path.join(os.path.dirname(__file__), "performance_text", f"{module_name}.py") module_path = os.path.join(
os.path.dirname(__file__), "performance_tester", f"{module_name}.py"
)
spec = importlib.util.spec_from_file_location(module_name, module_path) spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec) module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module) spec.loader.exec_module(module)
@@ -26,17 +33,21 @@ async def load_and_execute_module(module_name):
else: else:
print(f"模块 {module_name} 中没有找到 main 函数。") print(f"模块 {module_name} 中没有找到 main 函数。")
def get_module_description(module_name): def get_module_description(module_name):
module_path = os.path.join(os.path.dirname(__file__), "performance_text", f"{module_name}.py") module_path = os.path.join(
os.path.dirname(__file__), "performance_tester", f"{module_name}.py"
)
spec = importlib.util.spec_from_file_location(module_name, module_path) spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec) module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module) spec.loader.exec_module(module)
return getattr(module, "description", "暂无描述") return getattr(module, "description", "暂无描述")
def main(): def main():
modules = list_performance_text_modules() modules = list_performance_tester_modules()
if not modules: if not modules:
print("performance_text 目录中没有可用的性能测试工具。") print("performance_tester 目录中没有可用的性能测试工具。")
return return
print("可用的性能测试工具:") print("可用的性能测试工具:")
@@ -53,5 +64,6 @@ def main():
except ValueError: except ValueError:
print("请输入有效的数字。") print("请输入有效的数字。")
if __name__ == "__main__": if __name__ == "__main__":
main() main()
@@ -7,14 +7,17 @@ from typing import Dict
import aiohttp import aiohttp
from tabulate import tabulate from tabulate import tabulate
from core.utils.asr import create_instance as create_stt_instance from core.utils.asr import create_instance as create_stt_instance
from config.settings import load_config
# 设置全局日志级别为WARNING,抑制INFO级别日志 # 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "语音识别模型性能测试" description = "语音识别模型性能测试"
class ASRPerformanceTester: class ASRPerformanceTester:
def __init__(self): def __init__(self):
self.config = self._load_config_from_data_dir() self.config = load_config()
self.test_wav_list = self._load_test_wav_files() self.test_wav_list = self._load_test_wav_files()
self.results = {"stt": {}} self.results = {"stt": {}}
@@ -22,29 +25,6 @@ class ASRPerformanceTester:
print(f"[DEBUG] 加载的ASR配置: {self.config.get('ASR', {})}") print(f"[DEBUG] 加载的ASR配置: {self.config.get('ASR', {})}")
print(f"[DEBUG] 音频文件数量: {len(self.test_wav_list)}") print(f"[DEBUG] 音频文件数量: {len(self.test_wav_list)}")
def _load_config_from_data_dir(self) -> Dict:
"""从 data 目录加载所有 .config.yaml 文件的配置"""
config = {"ASR": {}}
data_dir = os.path.join(os.getcwd(), "data")
print(f"[DEBUG] 扫描配置文件目录: {data_dir}")
for root, _, files in os.walk(data_dir):
for file in files:
if file.endswith(".config.yaml"):
file_path = os.path.join(root, file)
try:
with open(file_path, "r", encoding="utf-8") as f:
import yaml
file_config = yaml.safe_load(f)
# 兼容大小写的 ASR/asr 配置
asr_config = file_config.get("ASR") or file_config.get("asr")
if asr_config:
config["ASR"].update(asr_config)
print(f"[DEBUG] 从 {file_path} 加载 ASR 配置成功")
except Exception as e:
print(f" 加载配置文件 {file_path} 失败: {str(e)}")
return config
def _load_test_wav_files(self) -> list: def _load_test_wav_files(self) -> list:
"""加载测试用的音频文件(添加路径调试)""" """加载测试用的音频文件(添加路径调试)"""
wav_root = os.path.join(os.getcwd(), "config", "assets") wav_root = os.path.join(os.getcwd(), "config", "assets")
@@ -8,14 +8,17 @@ import yaml
import aiohttp import aiohttp
from tabulate import tabulate from tabulate import tabulate
from core.utils.llm import create_instance as create_llm_instance from core.utils.llm import create_instance as create_llm_instance
from config.settings import load_config
# 设置全局日志级别为 WARNING,抑制 INFO 级别日志 # 设置全局日志级别为 WARNING,抑制 INFO 级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "大语言模型性能测试" description = "大语言模型性能测试"
class LLMPerformanceTester: class LLMPerformanceTester:
def __init__(self): def __init__(self):
self.config = self._load_config() self.config = load_config()
self.test_sentences = self.config.get("module_test", {}).get( self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences", "test_sentences",
[ [
@@ -26,24 +29,6 @@ class LLMPerformanceTester:
) )
self.results = {} self.results = {}
def _load_config(self) -> Dict:
"""从 data/.config.yaml 加载配置"""
config = {}
config_file_path = os.path.join(os.getcwd(), "data", ".config.yaml")
print(f"[DEBUG] 加载配置文件: {config_file_path}")
if not os.path.exists(config_file_path):
print(f"[DEBUG] 配置文件 {config_file_path} 不存在")
return config
try:
with open(config_file_path, "r", encoding="utf-8") as f:
config = yaml.safe_load(f)
print(f"[DEBUG] 从 {config_file_path} 加载配置成功")
except Exception as e:
print(f"加载配置文件 {config_file_path} 失败: {str(e)}")
return config
async def _check_ollama_service(self, base_url: str, model_name: str) -> bool: async def _check_ollama_service(self, base_url: str, model_name: str) -> bool:
"""异步检查 Ollama 服务状态""" """异步检查 Ollama 服务状态"""
async with aiohttp.ClientSession() as session: async with aiohttp.ClientSession() as session:
@@ -8,14 +8,17 @@ from tabulate import tabulate
# 确保从 core.utils.tts 导入 create_tts_instance # 确保从 core.utils.tts 导入 create_tts_instance
from core.utils.tts import create_instance as create_tts_instance from core.utils.tts import create_instance as create_tts_instance
from config.settings import load_config
# 设置全局日志级别为 WARNING # 设置全局日志级别为 WARNING
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "非流式语音合成性能测试" description = "非流式语音合成性能测试"
class TTSPerformanceTester: class TTSPerformanceTester:
def __init__(self): def __init__(self):
self.config = self._load_config_from_data_dir() self.config = load_config()
self.test_sentences = self.config.get("module_test", {}).get( self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences", "test_sentences",
[ [
@@ -26,27 +29,6 @@ class TTSPerformanceTester:
) )
self.results = {} self.results = {}
def _load_config_from_data_dir(self) -> Dict:
"""从 data 目录加载所有 .config.yaml 文件的配置"""
config = {"TTS": {}}
data_dir = os.path.join(os.getcwd(), "data")
print(f"[DEBUG] 扫描配置文件目录: {data_dir}")
for root, _, files in os.walk(data_dir):
for file in files:
if file.endswith(".config.yaml"):
file_path = os.path.join(root, file)
try:
with open(file_path, "r", encoding="utf-8") as f:
file_config = yaml.safe_load(f)
tts_config = file_config.get("TTS")
if tts_config:
config["TTS"].update(tts_config)
print(f"[DEBUG] 从 {file_path} 加载 TTS 配置成功")
except Exception as e:
print(f"加载配置文件 {file_path} 失败: {str(e)}")
return config
async def _test_tts(self, tts_name: str, config: Dict) -> Dict: async def _test_tts(self, tts_name: str, config: Dict) -> Dict:
"""测试单个TTS模块的性能""" """测试单个TTS模块的性能"""
try: try:
@@ -107,19 +89,19 @@ class TTSPerformanceTester:
table = [] table = []
for name, data in self.results.items(): for name, data in self.results.items():
if data["errors"] == 0: if data["errors"] == 0:
table.append([ table.append(
name, [name, f"{data['avg_time']:.3f}秒/句", len(self.test_sentences[:3])]
f"{data['avg_time']:.3f}秒/句", )
len(self.test_sentences[:3])
])
print("\nTTS性能测试结果:") print("\nTTS性能测试结果:")
print(tabulate( print(
table, tabulate(
headers=["TTS模块", "平均耗时", "测试句子数"], table,
tablefmt="github", headers=["TTS模块", "平均耗时", "测试句子数"],
colalign=("left", "right", "right") tablefmt="github",
)) colalign=("left", "right", "right"),
)
)
async def run(self): async def run(self):
"""执行测试""" """执行测试"""
@@ -145,11 +127,13 @@ class TTSPerformanceTester:
# 打印结果 # 打印结果
self._print_results() self._print_results()
#为了performance_tester.py的调用需求
# 为了performance_tester.py的调用需求
async def main(): async def main():
tester = TTSPerformanceTester() tester = TTSPerformanceTester()
await tester.run() await tester.run()
if __name__ == "__main__": if __name__ == "__main__":
tester = TTSPerformanceTester() tester = TTSPerformanceTester()
asyncio.run(tester.run()) asyncio.run(tester.run())
@@ -3,24 +3,20 @@ import asyncio
import logging import logging
import statistics import statistics
import base64 import base64
import yaml
from typing import Dict from typing import Dict
from tabulate import tabulate from tabulate import tabulate
from core.utils.vllm import create_instance from core.utils.vllm import create_instance
from config.settings import load_config
# 设置全局日志级别为WARNING,抑制INFO级别日志 # 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "视觉识别模型性能测试" description = "视觉识别模型性能测试"
class AsyncVisionPerformanceTester: class AsyncVisionPerformanceTester:
def __init__(self): def __init__(self):
# 从data/.config.yaml读取配置 self.config = load_config()
config_path = os.path.join("data", ".config.yaml")
if not os.path.exists(config_path):
raise FileNotFoundError(f"配置文件 {config_path} 不存在")
with open(config_path, "r", encoding="utf-8") as f:
self.config = yaml.safe_load(f) or {}
self.test_images = [ self.test_images = [
"../../docs/images/demo1.png", "../../docs/images/demo1.png",
@@ -1,651 +0,0 @@
import asyncio
import logging
import os
import statistics
import time
from typing import Dict
import yaml
import aiohttp
from tabulate import tabulate
from core.utils.asr import create_instance as create_stt_instance
from core.utils.llm import create_instance as create_llm_instance
from core.utils.tts import create_instance as create_tts_instance
# 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING)
description = "基础性能测试工具"
class AsyncPerformanceTester:
def __init__(self):
# 从data/.config.yaml读取配置
config_path = os.path.join("data", ".config.yaml")
if not os.path.exists(config_path):
raise FileNotFoundError(f"配置文件 {config_path} 不存在")
with open(config_path, "r", encoding="utf-8") as f:
self.config = yaml.safe_load(f) or {}
self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences",
[
"你好,请介绍一下你自己",
"What's the weather like today?",
"请用100字概括量子计算的基本原理和应用前景",
],
)
self.test_wav_list = []
self.wav_root = r"config/assets"
for file_name in os.listdir(self.wav_root):
file_path = os.path.join(self.wav_root, file_name)
# 检查文件大小是否大于300KB
if os.path.getsize(file_path) > 300 * 1024: # 300KB = 300 * 1024 bytes
with open(file_path, "rb") as f:
self.test_wav_list.append(f.read())
self.results = {"llm": {}, "tts": {}, "stt": {}, "combinations": []}
async def _check_ollama_service(self, base_url: str, model_name: str) -> bool:
"""异步检查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:
logging.getLogger("core.providers.tts.base").setLevel(logging.WARNING)
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 {"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_stt(self, stt_name: str, config: Dict) -> Dict:
"""异步测试单个STT性能"""
try:
logging.getLogger("core.providers.asr.base").setLevel(logging.WARNING)
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"⏭️ STT {stt_name} 未配置access_token/api_key,已跳过")
return {"name": stt_name, "type": "stt", "errors": 1}
module_type = config.get("type", stt_name)
stt = create_stt_instance(module_type, config, delete_audio_file=True)
stt.audio_format = "pcm"
print(f"🎵 测试 STT: {stt_name}")
text, _ = await stt.speech_to_text(
[self.test_wav_list[0]], "1", stt.audio_format
)
if text is None:
print(f"{stt_name} 连接失败")
return {"name": stt_name, "type": "stt", "errors": 1}
total_time = 0
test_count = len(self.test_wav_list)
for i, sentence in enumerate(self.test_wav_list, 1):
start = time.time()
text, _ = await stt.speech_to_text([sentence], "1", stt.audio_format)
duration = time.time() - start
total_time += duration
if text:
print(f"{stt_name} [{i}/{test_count}]")
else:
print(f"{stt_name} [{i}/{test_count}]")
return {"name": stt_name, "type": "stt", "errors": 1}
return {
"name": stt_name,
"type": "stt",
"avg_time": total_time / test_count,
"errors": 0,
}
except Exception as e:
print(f"⚠️ {stt_name} 测试失败: {str(e)}")
return {"name": stt_name, "type": "stt", "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)
llm = create_llm_instance(module_type, config)
# 统一使用UTF-8编码
test_sentences = [
s.encode("utf-8").decode("utf-8") for s in self.test_sentences
]
# 创建所有句子的测试任务
sentence_tasks = []
for sentence in test_sentences:
sentence_tasks.append(
self._test_single_sentence(llm_name, llm, sentence)
)
# 并发执行所有句子测试
sentence_results = await asyncio.gather(*sentence_tasks)
# 处理结果
valid_results = [r for r in sentence_results if r is not None]
if not valid_results:
print(f"⚠️ {llm_name} 无有效数据,可能配置错误")
return {"name": llm_name, "type": "llm", "errors": 1}
first_token_times = [r["first_token_time"] for r in valid_results]
response_times = [r["response_time"] for r in valid_results]
# 过滤异常数据
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:
print(f"⚠️ {llm_name} 有效数据不足,可能网络不稳定")
return {"name": llm_name, "type": "llm", "errors": 1}
return {
"name": llm_name,
"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_response": (
statistics.stdev(response_times) if len(response_times) > 1 else 0
),
"errors": 0,
}
except Exception as e:
print(f"LLM {llm_name} 测试失败: {str(e)}")
return {"name": llm_name, "type": "llm", "errors": 1}
async def _test_single_sentence(self, llm_name: str, llm, sentence: str) -> Dict:
"""测试单个句子的性能"""
try:
print(f"📝 {llm_name} 开始测试: {sentence[:20]}...")
sentence_start = time.time()
first_token_received = False
first_token_time = None
async def process_response():
nonlocal first_token_received, first_token_time
for chunk in llm.response(
"perf_test", [{"role": "user", "content": sentence}]
):
if not first_token_received and chunk.strip() != "":
first_token_time = time.time() - sentence_start
first_token_received = True
print(f"{llm_name} 首个Token: {first_token_time:.3f}s")
yield chunk
response_chunks = []
async for chunk in process_response():
response_chunks.append(chunk)
response_time = time.time() - sentence_start
print(f"{llm_name} 完成响应: {response_time:.3f}s")
if first_token_time is None:
first_token_time = (
response_time # 如果没有检测到first token,使用总响应时间
)
return {
"name": llm_name,
"type": "llm",
"first_token_time": first_token_time,
"response_time": response_time,
}
except Exception as e:
print(f"⚠️ {llm_name} 句子测试失败: {str(e)}")
return None
def _generate_combinations(self):
"""生成最佳组合建议"""
valid_llms = [
k
for k, v in self.results["llm"].items()
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_stt = [k for k, v in self.results["stt"].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
)
min_stt_time = (
min([self.results["stt"][stt]["avg_time"] for stt in valid_stt])
if valid_stt
else 1
)
for llm in valid_llms:
for tts in valid_tts:
for stt in valid_stt:
# 计算相对性能分数(越小越好)
llm_score = (
self.results["llm"][llm]["avg_first_token"] / min_first_token
)
tts_score = self.results["tts"][tts]["avg_time"] / min_tts_time
stt_score = self.results["stt"][stt]["avg_time"] / min_stt_time
# 计算稳定性分数(标准差/平均值,越小越稳定)
llm_stability = (
self.results["llm"][llm]["std_first_token"]
/ self.results["llm"][llm]["avg_first_token"]
)
# 综合得分(考虑性能和稳定性)
# LLM得分: 性能权重(70%) + 稳定性权重(30%)
llm_final_score = llm_score * 0.7 + llm_stability * 0.3
# 总分 = LLM得分(70%) + TTS得分(30%) + STT得分(30%)
total_score = (
llm_final_score * 0.7 + tts_score * 0.3 + stt_score * 0.3
)
self.results["combinations"].append(
{
"llm": llm,
"tts": tts,
"stt": stt,
"score": total_score,
"details": {
"llm_first_token": self.results["llm"][llm][
"avg_first_token"
],
"llm_stability": llm_stability,
"tts_time": self.results["tts"][tts]["avg_time"],
"stt_time": self.results["stt"][stt]["avg_time"],
},
}
)
# 分数越小越好
self.results["combinations"].sort(key=lambda x: x["score"])
def _print_results(self):
"""打印测试结果"""
llm_table = []
for name, data in self.results["llm"].items():
if data["errors"] == 0:
stability = data["std_first_token"] / data["avg_first_token"]
llm_table.append(
[
name, # 不需要固定宽度,让tabulate自己处理对齐
f"{data['avg_first_token']:.3f}",
f"{data['avg_response']:.3f}",
f"{stability:.3f}",
]
)
if llm_table:
print("\nLLM 性能排行:\n")
print(
tabulate(
llm_table,
headers=["模型名称", "首字耗时", "总耗时", "稳定性"],
tablefmt="github",
colalign=("left", "right", "right", "right"),
disable_numparse=True,
)
)
else:
print("\n⚠️ 没有可用的LLM模块进行测试。")
tts_table = []
for name, data in self.results["tts"].items():
if data["errors"] == 0:
tts_table.append([name, f"{data['avg_time']:.3f}"]) # 不需要固定宽度
if tts_table:
print("\nTTS 性能排行:\n")
print(
tabulate(
tts_table,
headers=["模型名称", "合成耗时"],
tablefmt="github",
colalign=("left", "right"),
disable_numparse=True,
)
)
else:
print("\n⚠️ 没有可用的TTS模块进行测试。")
stt_table = []
for name, data in self.results["stt"].items():
if data["errors"] == 0:
stt_table.append([name, f"{data['avg_time']:.3f}"]) # 不需要固定宽度
if stt_table:
print("\nSTT 性能排行:\n")
print(
tabulate(
stt_table,
headers=["模型名称", "合成耗时"],
tablefmt="github",
colalign=("left", "right"),
disable_numparse=True,
)
)
else:
print("\n⚠️ 没有可用的STT模块进行测试。")
if self.results["combinations"]:
print("\n推荐配置组合 (得分越小越好):\n")
combo_table = []
for combo in self.results["combinations"][:]:
combo_table.append(
[
f"{combo['llm']} + {combo['tts']} + {combo['stt']}", # 不需要固定宽度
f"{combo['score']:.3f}",
f"{combo['details']['llm_first_token']:.3f}",
f"{combo['details']['llm_stability']:.3f}",
f"{combo['details']['tts_time']:.3f}",
f"{combo['details']['stt_time']:.3f}",
]
)
print(
tabulate(
combo_table,
headers=[
"组合方案",
"综合得分",
"LLM首字耗时",
"稳定性",
"TTS合成耗时",
"STT合成耗时",
],
tablefmt="github",
colalign=("left", "right", "right", "right", "right", "right"),
disable_numparse=True,
)
)
else:
print("\n⚠️ 没有可用的模块组合建议。")
def _process_results(self, all_results):
"""处理测试结果"""
for result in all_results:
if result["errors"] == 0:
if result["type"] == "llm":
self.results["llm"][result["name"]] = result
elif result["type"] == "tts":
self.results["tts"][result["name"]] = result
elif result["type"] == "stt":
self.results["stt"][result["name"]] = result
else:
pass
async def run(self):
"""执行全量异步测试"""
print("🔍 开始筛选可用模块...")
# 创建所有测试任务
all_tasks = []
# LLM测试任务
if self.config.get("LLM") is not None:
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测试任务
if self.config.get("TTS") is not None:
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))
# STT测试任务
if len(self.test_wav_list) >= 1:
if self.config.get("ASR") is not None:
for stt_name, config in self.config.get("ASR", {}).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"⏭️ ASR {stt_name} 未配置access_token/api_key,已跳过")
continue
print(f"🎵 添加ASR测试任务: {stt_name}")
all_tasks.append(self._test_stt(stt_name, config))
else:
print(f"\n⚠️ {self.wav_root} 路径下没有音频文件,已跳过STT测试任务")
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(
f"✅ 找到 {len([t for t in all_tasks if '_test_stt' in str(t)])} 个可用STT模块"
)
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
# 处理STT结果
for result in [
r
for r in all_results
if r and isinstance(r, dict) and r.get("type") == "stt"
]:
if result["errors"] == 0:
self.results["stt"][result["name"]] = result
# 生成组合建议并打印结果
print("\n📊 生成测试报告...")
self._generate_combinations()
self._print_results()
async def main():
tester = AsyncPerformanceTester()
await tester.run()
if __name__ == "__main__":
asyncio.run(main())