From 768d2b82d6441bc312c50f6abc38268625862d89 Mon Sep 17 00:00:00 2001 From: FAN-yeB <1442100690@qq.com> Date: Thu, 14 Aug 2025 10:01:43 +0800 Subject: [PATCH] =?UTF-8?q?update:=E6=9B=B4=E6=96=B0=E6=80=A7=E8=83=BD?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E5=B7=A5=E5=85=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/performance_tester.md | 29 ++ main/xiaozhi-server/performance_test_tool.py | 56 ++++ .../performance_tester.py | 2 +- .../performance_tester_asr.py | 169 +++++++++++ .../performance_tester_llm.py | 264 ++++++++++++++++++ .../performance_tester_tts.py | 155 ++++++++++ .../performance_tester_vllm.py | 2 +- 7 files changed, 675 insertions(+), 2 deletions(-) create mode 100644 docs/performance_tester.md create mode 100644 main/xiaozhi-server/performance_test_tool.py rename main/xiaozhi-server/{ => performance_text}/performance_tester.py (99%) create mode 100644 main/xiaozhi-server/performance_text/performance_tester_asr.py create mode 100644 main/xiaozhi-server/performance_text/performance_tester_llm.py create mode 100644 main/xiaozhi-server/performance_text/performance_tester_tts.py rename main/xiaozhi-server/{ => performance_text}/performance_tester_vllm.py (99%) diff --git a/docs/performance_tester.md b/docs/performance_tester.md new file mode 100644 index 00000000..97cf8945 --- /dev/null +++ b/docs/performance_tester.md @@ -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.得到测试结果 \ No newline at end of file diff --git a/main/xiaozhi-server/performance_test_tool.py b/main/xiaozhi-server/performance_test_tool.py new file mode 100644 index 00000000..0ed85d65 --- /dev/null +++ b/main/xiaozhi-server/performance_test_tool.py @@ -0,0 +1,56 @@ +import os +import importlib.util +import asyncio + +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() diff --git a/main/xiaozhi-server/performance_tester.py b/main/xiaozhi-server/performance_text/performance_tester.py similarity index 99% rename from main/xiaozhi-server/performance_tester.py rename to main/xiaozhi-server/performance_text/performance_tester.py index a5fc5221..5d6af469 100644 --- a/main/xiaozhi-server/performance_tester.py +++ b/main/xiaozhi-server/performance_text/performance_tester.py @@ -15,7 +15,7 @@ from core.utils.tts import create_instance as create_tts_instance # 设置全局日志级别为WARNING,抑制INFO级别日志 logging.basicConfig(level=logging.WARNING) - +description = "基础性能测试工具" class AsyncPerformanceTester: def __init__(self): diff --git a/main/xiaozhi-server/performance_text/performance_tester_asr.py b/main/xiaozhi-server/performance_text/performance_tester_asr.py new file mode 100644 index 00000000..044b8da5 --- /dev/null +++ b/main/xiaozhi-server/performance_text/performance_tester_asr.py @@ -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()) diff --git a/main/xiaozhi-server/performance_text/performance_tester_llm.py b/main/xiaozhi-server/performance_text/performance_tester_llm.py new file mode 100644 index 00000000..0d5d9562 --- /dev/null +++ b/main/xiaozhi-server/performance_text/performance_tester_llm.py @@ -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()) diff --git a/main/xiaozhi-server/performance_text/performance_tester_tts.py b/main/xiaozhi-server/performance_text/performance_tester_tts.py new file mode 100644 index 00000000..eab2cd4b --- /dev/null +++ b/main/xiaozhi-server/performance_text/performance_tester_tts.py @@ -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()) diff --git a/main/xiaozhi-server/performance_tester_vllm.py b/main/xiaozhi-server/performance_text/performance_tester_vllm.py similarity index 99% rename from main/xiaozhi-server/performance_tester_vllm.py rename to main/xiaozhi-server/performance_text/performance_tester_vllm.py index 4dafddcc..469569b5 100644 --- a/main/xiaozhi-server/performance_tester_vllm.py +++ b/main/xiaozhi-server/performance_text/performance_tester_vllm.py @@ -10,7 +10,7 @@ from core.utils.vllm import create_instance # 设置全局日志级别为WARNING,抑制INFO级别日志 logging.basicConfig(level=logging.WARNING) - +description = "视觉识别模型性能测试" class AsyncVisionPerformanceTester: def __init__(self):