update:合并非tts代码

This commit is contained in:
hrz
2025-05-21 13:18:12 +08:00
parent ede2bc6a4e
commit 851365fb58
65 changed files with 5423 additions and 1943 deletions
+362 -156
View File
@@ -1,16 +1,17 @@
import time
import aiohttp
import asyncio
import logging
import os
import statistics
import time
from typing import Dict
import aiohttp
from tabulate import tabulate
from typing import Dict, List
from config.settings import load_config
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
from core.utils.util import read_config
import statistics
from config.settings import get_config_file
import inspect
import os
import logging
# 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING)
@@ -18,17 +19,26 @@ logging.basicConfig(level=logging.WARNING)
class AsyncPerformanceTester:
def __init__(self):
self.config = read_config(get_config_file())
self.config = load_config()
self.test_sentences = self.config.get("module_test", {}).get(
"test_sentences",
["你好,请介绍一下你自己", "What's the weather like today?",
"请用100字概括量子计算的基本原理和应用前景"]
[
"你好,请介绍一下你自己",
"What's the weather like today?",
"请用100字概括量子计算的基本原理和应用前景",
],
)
self.results = {
"llm": {},
"tts": {},
"combinations": []
}
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服务状态"""
@@ -46,7 +56,9 @@ class AsyncPerformanceTester:
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} 下载")
print(
f"🚫 Ollama模型 {model_name} 未找到,请先使用 ollama pull {model_name} 下载"
)
return False
else:
print(f"🚫 无法获取Ollama模型列表")
@@ -62,17 +74,16 @@ class AsyncPerformanceTester:
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):
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
)
module_type = config.get("type", tts_name)
tts = create_tts_instance(module_type, config, delete_audio_file=True)
print(f"🎵 测试 TTS: {tts_name}")
@@ -103,20 +114,71 @@ class AsyncPerformanceTester:
"name": tts_name,
"type": "tts",
"avg_time": total_time / test_count,
"errors": 0
"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")
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")
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')
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}
@@ -124,21 +186,27 @@ class AsyncPerformanceTester:
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"]):
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)
# 统一使用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
]
# 创建所有句子的测试任务
sentence_tasks = []
for sentence in test_sentences:
sentence_tasks.append(self._test_single_sentence(llm_name, llm, sentence))
sentence_tasks.append(
self._test_single_sentence(llm_name, llm, sentence)
)
# 并发执行所有句子测试
sentence_results = await asyncio.gather(*sentence_tasks)
@@ -166,9 +234,15 @@ class AsyncPerformanceTester:
"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
"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)}")
@@ -184,8 +258,10 @@ class AsyncPerformanceTester:
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() != '':
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")
@@ -199,13 +275,15 @@ class AsyncPerformanceTester:
print(f"{llm_name} 完成响应: {response_time:.3f}s")
if first_token_time is None:
first_token_time = response_time # 如果没有检测到first token,使用总响应时间
first_token_time = (
response_time # 如果没有检测到first token,使用总响应时间
)
return {
"name": llm_name,
"type": "llm",
"first_token_time": first_token_time,
"response_time": response_time
"response_time": response_time,
}
except Exception as e:
print(f"⚠️ {llm_name} 句子测试失败: {str(e)}")
@@ -214,42 +292,71 @@ class AsyncPerformanceTester:
def _generate_combinations(self):
"""生成最佳组合建议"""
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
]
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_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:
# 计算相对性能分数(越小越好)
llm_score = self.results["llm"][llm]["avg_first_token"] / min_first_token
tts_score = self.results["tts"][tts]["avg_time"] / min_tts_time
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_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%) + 稳定性权重(30%)
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
# 总分 = 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,
"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"]
}
})
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"])
@@ -260,64 +367,98 @@ class AsyncPerformanceTester:
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}"
])
llm_table.append(
[
name, # 不需要固定宽度,让tabulate自己处理对齐
f"{data['avg_first_token']:.3f}",
f"{data['avg_response']:.3f}",
f"{stability:.3f}",
]
)
if llm_table:
print("\nLLM 性能排行:")
print(tabulate(
llm_table,
headers=["模型名称", "首字耗时", "总耗时", "稳定性"],
tablefmt="github",
colalign=("left", "right", "right", "right"),
disable_numparse=True
))
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}"
])
tts_table.append([name, f"{data['avg_time']:.3f}"]) # 不需要固定宽度
if tts_table:
print("\nTTS 性能排行:")
print(tabulate(
tts_table,
headers=["模型名称", "合成耗时"],
tablefmt="github",
colalign=("left", "right"),
disable_numparse=True
))
print("\nTTS 性能排行:\n")
print(
tabulate(
tts_table,
headers=["模型名称", "合成耗时"],
tablefmt="github",
colalign=("left", "right"),
disable_numparse=True,
)
)
else:
print("\n⚠️ 没有可用的TTS模块进行测试。")
if self.results["combinations"]:
print("\n推荐配置组合 (得分越小越好):")
combo_table = []
for combo in self.results["combinations"][:5]:
combo_table.append([
f"{combo['llm']} + {combo['tts']}", # 不需要固定宽度
f"{combo['score']:.3f}",
f"{combo['details']['llm_first_token']:.3f}",
f"{combo['details']['llm_stability']:.3f}",
f"{combo['details']['tts_time']:.3f}"
])
stt_table = []
for name, data in self.results["stt"].items():
if data["errors"] == 0:
stt_table.append([name, f"{data['avg_time']:.3f}"]) # 不需要固定宽度
print(tabulate(
combo_table,
headers=["组合方案", "综合得分", "LLM首字耗时", "稳定性", "TTS合成耗时"],
tablefmt="github",
colalign=("left", "right", "right", "right", "right"),
disable_numparse=True
))
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⚠️ 没有可用的模块组合建议。")
@@ -327,8 +468,12 @@ class AsyncPerformanceTester:
if result["errors"] == 0:
if result["type"] == "llm":
self.results["llm"][result["name"]] = result
else:
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):
"""执行全量异步测试"""
@@ -338,50 +483,83 @@ class AsyncPerformanceTester:
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")
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
if not await self._check_ollama_service(base_url, model_name):
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
print(f"📋 添加LLM测试任务: {llm_name}")
module_type = config.get('type', llm_name)
llm = create_llm_instance(module_type, config)
if not await self._check_ollama_service(base_url, model_name):
continue
# 为每个句子创建独立任务
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))
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))
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模块")
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")
# 并发执行所有测试任务
@@ -389,7 +567,11 @@ class AsyncPerformanceTester:
# 处理LLM结果
llm_results = {}
for result in [r for r in all_results if r and isinstance(r, dict) and r.get("type") == "llm"]:
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] = {
@@ -397,9 +579,11 @@ class AsyncPerformanceTester:
"type": "llm",
"first_token_times": [],
"response_times": [],
"errors": 0
"errors": 0,
}
llm_results[llm_name]["first_token_times"].append(result["first_token_time"])
llm_results[llm_name]["first_token_times"].append(
result["first_token_time"]
)
llm_results[llm_name]["response_times"].append(result["response_time"])
# 计算LLM平均值和标准差
@@ -408,19 +592,41 @@ class AsyncPerformanceTester:
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
"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"]:
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()
@@ -433,4 +639,4 @@ async def main():
if __name__ == "__main__":
asyncio.run(main())
asyncio.run(main())