From 5360cbf19c44fa285d726dd730d0b72787a74184 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Tue, 20 May 2025 11:14:20 +0800 Subject: [PATCH] =?UTF-8?q?update=EF=BC=9A=E4=BC=98=E5=8C=96asr=E6=B5=8B?= =?UTF-8?q?=E9=80=9F=E6=BA=90=EF=BC=9A=E4=BB=8E=E5=B7=B2=E6=9C=89=E7=9A=84?= =?UTF-8?q?=E5=87=A0=E4=B8=AA=E6=96=87=E4=BB=B6=E4=B8=AD=E6=B5=8B=E9=80=9F?= =?UTF-8?q?=E5=BA=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/config.yaml | 2 +- main/xiaozhi-server/performance_tester.py | 56 +++++++++++++---------- 2 files changed, 32 insertions(+), 26 deletions(-) diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml index c1ca2dbb..05756cd7 100644 --- a/main/xiaozhi-server/config.yaml +++ b/main/xiaozhi-server/config.yaml @@ -411,7 +411,7 @@ LLM: base_url: https://host/api/v1 # 你可以在这里找到你的api_key # https://cloud.tryfastgpt.ai/account/apikey - api_key: fastgpt-xxx + api_key: 你的fastgpt密钥 variables: k: "v" k2: "v2" diff --git a/main/xiaozhi-server/performance_tester.py b/main/xiaozhi-server/performance_tester.py index fbd460c7..80503166 100644 --- a/main/xiaozhi-server/performance_tester.py +++ b/main/xiaozhi-server/performance_tester.py @@ -30,11 +30,13 @@ class AsyncPerformanceTester: ) self.test_wav_list = [] - self.wav_root = r".\config\assets\test_asr" - os.makedirs(self.wav_root, exist_ok=True) - for wav_name in os.listdir(self.wav_root): - with open(os.path.join(self.wav_root, wav_name), "rb") as f: - self.test_wav_list.append(f.read()) + 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": []} @@ -73,9 +75,9 @@ class AsyncPerformanceTester: 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 + 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} @@ -185,7 +187,7 @@ class AsyncPerformanceTester: 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"] + x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"] ): print(f"🚫 跳过未配置的LLM: {llm_name}") return {"name": llm_name, "type": "llm", "errors": 1} @@ -257,7 +259,7 @@ class AsyncPerformanceTester: async def process_response(): nonlocal first_token_received, first_token_time for chunk in llm.response( - "perf_test", [{"role": "user", "content": sentence}] + "perf_test", [{"role": "user", "content": sentence}] ): if not first_token_received and chunk.strip() != "": first_token_time = time.time() - sentence_start @@ -319,15 +321,15 @@ class AsyncPerformanceTester: for stt in valid_stt: # 计算相对性能分数(越小越好) llm_score = ( - self.results["llm"][llm]["avg_first_token"] / min_first_token + 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"] + self.results["llm"][llm]["std_first_token"] + / self.results["llm"][llm]["avg_first_token"] ) # 综合得分(考虑性能和稳定性) @@ -335,7 +337,9 @@ class AsyncPerformanceTester: 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 + total_score = ( + llm_final_score * 0.7 + tts_score * 0.3 + stt_score * 0.3 + ) self.results["combinations"].append( { @@ -484,12 +488,12 @@ class AsyncPerformanceTester: # 检查配置有效性 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 ["你的"] + 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"] + x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"] ): print(f"⏭️ LLM {llm_name} 未配置api_key,已跳过") continue @@ -512,16 +516,18 @@ class AsyncPerformanceTester: # 为每个句子创建独立任务 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)) + 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 + 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 @@ -534,9 +540,9 @@ class AsyncPerformanceTester: 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 + 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 @@ -587,9 +593,9 @@ class AsyncPerformanceTester: "name": llm_name, "type": "llm", "avg_response": sum(data["response_times"]) - / len(data["response_times"]), + / len(data["response_times"]), "avg_first_token": sum(data["first_token_times"]) - / len(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