Merge pull request #2035 from xinnan-tech/performance_text

update:更新各组件性能测试功能
This commit is contained in:
CGD
2025-08-14 17:58:37 +08:00
committed by GitHub
7 changed files with 694 additions and 6 deletions
+29
View File
@@ -0,0 +1,29 @@
# 语音识别、大语言模型、非流式语音合成、视觉模型的性能测试工具使用指南
1.在main/xiaozhi-server目录下创建data目录
2.在data目录下创建.config.yaml文件
3.在.data/config.yaml中,写入你的语音识别、大语言模型、非流式语音合成、视觉模型的参数
例如:
```
LLM:
ChatGLMLLM:
# 定义LLM API类型
type: openai
# glm-4-flash 是免费的,但是还是需要注册填写api_key的
# 可在这里找到你的api key https://bigmodel.cn/usercenter/proj-mgmt/apikeys
model_name: glm-4-flash
url: https://open.bigmodel.cn/api/paas/v4/
api_key: 你的chat-glm web key
TTS:
VLLM:
ASR:
```
4.在main/xiaozhi-server目录下运行performance_text_tool.py:
```
python performance_test_tool.py
```
5.性能测试结果将保存在test_result目录。
6.得到测试结果
@@ -0,0 +1,57 @@
import os
import importlib.util
import asyncio
print("使用前请根据doc/performance_texter.md的说明准备配置。")
def list_performance_text_modules():
performance_text_dir = os.path.join(os.path.dirname(__file__), "performance_text")
modules = []
for file in os.listdir(performance_text_dir):
if file.endswith(".py"):
modules.append(file[:-3])
return modules
async def load_and_execute_module(module_name):
module_path = os.path.join(os.path.dirname(__file__), "performance_text", f"{module_name}.py")
spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
if hasattr(module, "main"):
main_func = module.main
if asyncio.iscoroutinefunction(main_func):
await main_func()
else:
main_func()
else:
print(f"模块 {module_name} 中没有找到 main 函数。")
def get_module_description(module_name):
module_path = os.path.join(os.path.dirname(__file__), "performance_text", f"{module_name}.py")
spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return getattr(module, "description", "暂无描述")
def main():
modules = list_performance_text_modules()
if not modules:
print("performance_text 目录中没有可用的性能测试工具。")
return
print("可用的性能测试工具:")
for idx, module in enumerate(modules, 1):
description = get_module_description(module)
print(f"{idx}. {module} - {description}")
try:
choice = int(input("请选择要调用的性能测试工具编号:")) - 1
if 0 <= choice < len(modules):
asyncio.run(load_and_execute_module(modules[choice]))
else:
print("无效的选择。")
except ValueError:
print("请输入有效的数字。")
if __name__ == "__main__":
main()
@@ -4,22 +4,29 @@ import os
import statistics import statistics
import time import time
from typing import Dict from typing import Dict
import yaml
import aiohttp import aiohttp
from tabulate import tabulate from tabulate import tabulate
from config.settings import load_config
from core.utils.asr import create_instance as create_stt_instance 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.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
# 设置全局日志级别为WARNING,抑制INFO级别日志 # 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "基础性能测试工具"
class AsyncPerformanceTester: class AsyncPerformanceTester:
def __init__(self): def __init__(self):
self.config = load_config() # 从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( self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences", "test_sentences",
[ [
@@ -0,0 +1,169 @@
import asyncio
import logging
import os
import time
from typing import Dict
import aiohttp
from tabulate import tabulate
from core.utils.asr import create_instance as create_stt_instance
# 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING)
description = "语音识别模型性能测试"
class ASRPerformanceTester:
def __init__(self):
self.config = self._load_config_from_data_dir()
self.test_wav_list = self._load_test_wav_files()
self.results = {"stt": {}}
# 调试日志
print(f"[DEBUG] 加载的ASR配置: {self.config.get('ASR', {})}")
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:
"""加载测试用的音频文件(添加路径调试)"""
wav_root = os.path.join(os.getcwd(), "config", "assets")
print(f"[DEBUG] 音频文件目录: {wav_root}")
test_wav_list = []
if os.path.exists(wav_root):
file_list = os.listdir(wav_root)
print(f"[DEBUG] 找到音频文件: {file_list}")
for file_name in file_list:
file_path = os.path.join(wav_root, file_name)
if os.path.getsize(file_path) > 300 * 1024: # 300KB
with open(file_path, "rb") as f:
test_wav_list.append(f.read())
else:
print(f" 目录不存在: {wav_root}")
return test_wav_list
async def _test_stt(self, stt_name: str, config: Dict) -> Dict:
"""异步测试单个STT性能(跳过无效配置)"""
try:
token_fields = ["access_token", "api_key", "token"]
# 忽略值为 "none" 的情况(需根据实际需求调整)
if any(
field in config
and str(config[field]).lower() 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
print(f" {stt_name} [{i}/{test_count}] 耗时: {duration:.2f}s")
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}
def _print_results(self):
"""打印测试结果"""
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("\nASR 性能排行:\n")
print(
tabulate(
stt_table,
headers=["模型名称", "平均耗时"],
tablefmt="github",
colalign=("left", "right"),
)
)
else:
print("\n 没有可用的ASR模块进行测试。")
async def run(self):
"""执行全量异步测试"""
print("开始筛选可用ASR模块...")
if not self.config.get("ASR"):
print("配置中未找到 ASR 模块")
return
all_tasks = []
for stt_name, config in self.config["ASR"].items():
print(f"[DEBUG] 检查 ASR 模块: {stt_name}, 配置: {config}")
all_tasks.append(self._test_stt(stt_name, config))
if not all_tasks:
print("没有可用的ASR模块进行测试。")
return
print("\n开始并发测试所有ASR模块...")
all_results = await asyncio.gather(*all_tasks, return_exceptions=True)
# 处理结果
for result in all_results:
if isinstance(result, dict) and result.get("type") == "stt":
if result["errors"] == 0:
self.results["stt"][result["name"]] = result
# 打印结果
print("\n测试完成")
self._print_results()
async def main():
tester = ASRPerformanceTester()
await tester.run()
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,264 @@
import asyncio
import logging
import os
import statistics
import time
from typing import Dict, Optional
import yaml
import aiohttp
from tabulate import tabulate
from core.utils.llm import create_instance as create_llm_instance
# 设置全局日志级别为 WARNING,抑制 INFO 级别日志
logging.basicConfig(level=logging.WARNING)
description = "大语言模型性能测试"
class LLMPerformanceTester:
def __init__(self):
self.config = self._load_config()
self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences",
[
"你好,请介绍一下你自己",
"What's the weather like today?",
"请用100字概括量子计算的基本原理和应用前景",
],
)
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:
"""异步检查 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("无法获取 Ollama 模型列表")
return False
return True
except Exception as e:
print(f"无法连接到 Ollama 服务: {str(e)}")
return False
async def _test_single_sentence(
self, llm_name: str, llm, sentence: str
) -> Optional[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")
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
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("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),
"errors": 0,
}
except Exception as e:
print(f"LLM {llm_name} 测试失败: {str(e)}")
return {"name": llm_name, "type": "llm", "errors": 1}
def _print_results(self):
"""打印测试结果"""
llm_table = []
for name, data in self.results.items():
if data["errors"] == 0:
llm_table.append(
[
name,
f"{data['avg_first_token']:.3f}",
f"{data['avg_response']:.3f}",
]
)
if llm_table:
print("\nLLM 性能排行:\n")
print(
tabulate(
llm_table,
headers=["模型名称", "首字耗时", "总耗时"],
tablefmt="github",
colalign=("left", "right", "right"),
disable_numparse=True,
)
)
else:
print("\n没有可用的 LLM 模块进行测试。")
async def run(self):
"""执行全量异步测试"""
print("开始筛选可用 LLM 模块...")
# 创建所有测试任务
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("Ollama 未配置 model_name")
continue
if not await self._check_ollama_service(base_url, model_name):
continue
print(f"添加 LLM 测试任务: {llm_name}")
all_tasks.append(self._test_llm(llm_name, config))
print(f"\n找到 {len(all_tasks)} 个可用 LLM 模块")
print("\n开始并发测试所有模块...\n")
# 并发执行所有测试任务
all_results = await asyncio.gather(*all_tasks, return_exceptions=True)
# 处理结果
for result in all_results:
if isinstance(result, dict) and result.get("errors") == 0:
self.results[result["name"]] = result
# 打印结果
print("\n生成测试报告...")
self._print_results()
async def main():
tester = LLMPerformanceTester()
await tester.run()
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,155 @@
import asyncio
import logging
import os
import time
from typing import Dict
import yaml
from tabulate import tabulate
# 确保从 core.utils.tts 导入 create_tts_instance
from core.utils.tts import create_instance as create_tts_instance
# 设置全局日志级别为 WARNING
logging.basicConfig(level=logging.WARNING)
description = "非流式语音合成性能测试"
class TTSPerformanceTester:
def __init__(self):
self.config = self._load_config_from_data_dir()
self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences",
[
"永和九年,岁在癸丑,暮春之初;",
"夫人之相与,俯仰一世,或取诸怀抱,悟言一室之内;或因寄所托,放浪形骸之外。虽趣舍万殊,静躁不同,",
"每览昔人兴感之由,若合一契,未尝不临文嗟悼,不能喻之于怀。固知一死生为虚诞,齐彭殇为妄作。",
],
)
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:
"""测试单个TTS模块的性能"""
try:
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, "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, "errors": 1}
total_time = 0
test_count = len(self.test_sentences[:3])
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, "errors": 1}
return {
"name": tts_name,
"avg_time": total_time / test_count,
"errors": 0,
}
except Exception as e:
print(f"{tts_name} 测试失败: {str(e)}")
return {"name": tts_name, "errors": 1}
def _print_results(self):
"""打印测试结果"""
if not self.results:
print("没有有效的TTS测试结果")
return
table = []
for name, data in self.results.items():
if data["errors"] == 0:
table.append([
name,
f"{data['avg_time']:.3f}秒/句",
len(self.test_sentences[:3])
])
print("\nTTS性能测试结果:")
print(tabulate(
table,
headers=["TTS模块", "平均耗时", "测试句子数"],
tablefmt="github",
colalign=("left", "right", "right")
))
async def run(self):
"""执行测试"""
print("开始TTS性能测试...")
if not self.config.get("TTS"):
print("配置文件中未找到TTS配置")
return
# 遍历所有TTS配置
tasks = []
for tts_name, config in self.config.get("TTS", {}).items():
tasks.append(self._test_tts(tts_name, config))
# 并发执行测试
results = await asyncio.gather(*tasks)
# 保存有效结果
for result in results:
if result["errors"] == 0:
self.results[result["name"]] = result
# 打印结果
self._print_results()
#为了performance_tester.py的调用需求
async def main():
tester = TTSPerformanceTester()
await tester.run()
if __name__ == "__main__":
tester = TTSPerformanceTester()
asyncio.run(tester.run())
@@ -3,18 +3,25 @@ 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 config.settings import load_config
from core.utils.vllm import create_instance from core.utils.vllm import create_instance
# 设置全局日志级别为WARNING,抑制INFO级别日志 # 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING) logging.basicConfig(level=logging.WARNING)
description = "视觉识别模型性能测试"
class AsyncVisionPerformanceTester: class AsyncVisionPerformanceTester:
def __init__(self): def __init__(self):
self.config = load_config() # 从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_images = [ self.test_images = [
"../../docs/images/demo1.png", "../../docs/images/demo1.png",
"../../docs/images/demo2.png", "../../docs/images/demo2.png",