mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-24 16:13:54 +08:00
Fix long audio bug (#158)
* update:异步生成音频 * update:优化LLM断句 --------- Co-authored-by: hrz <1710360675@qq.com>
This commit is contained in:
@@ -1,7 +1,5 @@
|
||||
import os
|
||||
import sys
|
||||
import asyncio
|
||||
from typing import List, Dict, Any
|
||||
|
||||
# 添加项目根目录到Python路径
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
@@ -10,10 +8,6 @@ sys.path.insert(0, project_root)
|
||||
|
||||
from config.logger import setup_logging
|
||||
import importlib
|
||||
from datetime import datetime
|
||||
from core.utils.util import is_segment
|
||||
from core.utils.util import get_string_no_punctuation_or_emoji
|
||||
from core.utils.util import read_config, get_project_dir
|
||||
|
||||
logger = setup_logging()
|
||||
|
||||
@@ -27,117 +21,3 @@ def create_instance(class_name, *args, **kwargs):
|
||||
return sys.modules[lib_name].LLMProvider(*args, **kwargs)
|
||||
|
||||
raise ValueError(f"不支持的LLM类型: {class_name},请检查该配置的type是否设置正确")
|
||||
|
||||
|
||||
async def test_single_model(llm_name: str, llm_config: Dict[str, Any], test_prompt: str, config: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""异步测试单个模型"""
|
||||
try:
|
||||
# 获取实际的LLM类型
|
||||
llm_type = llm_config["type"] if "type" in llm_config else llm_name
|
||||
llm = create_instance(llm_type, llm_config)
|
||||
|
||||
# 开始测试
|
||||
dialogue = []
|
||||
dialogue.append({"role": "system", "content": config.get("prompt")})
|
||||
dialogue.append({"role": "user", "content": test_prompt})
|
||||
|
||||
start_time = datetime.now()
|
||||
llm_responses = llm.response("test", dialogue)
|
||||
response_message = []
|
||||
first_response_time = None
|
||||
total_response_time = None
|
||||
start = 0
|
||||
full_response = ""
|
||||
|
||||
for content in llm_responses:
|
||||
response_message.append(content)
|
||||
full_response += content
|
||||
|
||||
if is_segment(response_message):
|
||||
segment_text = "".join(response_message[start:])
|
||||
segment_text = get_string_no_punctuation_or_emoji(segment_text)
|
||||
if len(segment_text) > 0:
|
||||
if first_response_time is None:
|
||||
first_response_time = (datetime.now() - start_time).total_seconds()
|
||||
start = len(response_message)
|
||||
|
||||
total_response_time = (datetime.now() - start_time).total_seconds()
|
||||
|
||||
return {
|
||||
"name": llm_name,
|
||||
"type": llm_type,
|
||||
"first_response_time": first_response_time,
|
||||
"total_response_time": total_response_time,
|
||||
"response_length": len(full_response),
|
||||
"status": "成功",
|
||||
"response": full_response
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
print(f"测试 {llm_name} 时发生错误: {str(e)}")
|
||||
return {
|
||||
"name": llm_name,
|
||||
"type": llm_config.get("type", llm_name),
|
||||
"first_response_time": None,
|
||||
"total_response_time": None,
|
||||
"response_length": 0,
|
||||
"status": f"失败 - {str(e)}",
|
||||
"response": ""
|
||||
}
|
||||
|
||||
|
||||
async def main():
|
||||
"""
|
||||
LLM模型响应速度测试和排行(异步版本)
|
||||
"""
|
||||
config = read_config(get_project_dir() + "config.yaml")
|
||||
test_prompt = "你好小智"
|
||||
|
||||
print("开始并发测试所有模型...")
|
||||
|
||||
# 创建所有模型的测试任务
|
||||
tasks = []
|
||||
for llm_name, llm_config in config["LLM"].items():
|
||||
task = asyncio.create_task(test_single_model(llm_name, llm_config, test_prompt, config))
|
||||
tasks.append(task)
|
||||
|
||||
# 等待所有测试完成
|
||||
test_results = await asyncio.gather(*tasks)
|
||||
|
||||
# 打印测试结果排行榜
|
||||
print("\n========= LLM模型性能测试排行榜 =========")
|
||||
print("测试提示词:", test_prompt)
|
||||
|
||||
# 过滤出成功的结果,并确保数值有效
|
||||
successful_results = [r for r in test_results if r["status"] == "成功" and r["first_response_time"] is not None]
|
||||
|
||||
if successful_results:
|
||||
print("\n1. 首次响应时间排行:")
|
||||
sorted_by_first = sorted(successful_results, key=lambda x: x["first_response_time"])
|
||||
for i, result in enumerate(sorted_by_first, 1):
|
||||
print(f"{i}. {result['name']}({result['type']}) - {result['first_response_time']:.2f}秒")
|
||||
print(f" 响应内容: {result['response'][:50]}...") # 只显示前50个字符
|
||||
|
||||
print("\n2. 总响应时间排行:")
|
||||
sorted_by_total = sorted(successful_results, key=lambda x: x["total_response_time"] or float('inf'))
|
||||
for i, result in enumerate(sorted_by_total, 1):
|
||||
if result["total_response_time"] is not None:
|
||||
print(f"{i}. {result['name']}({result['type']}) - {result['total_response_time']:.2f}秒")
|
||||
|
||||
print("\n3. 响应长度比较:")
|
||||
sorted_by_length = sorted(successful_results, key=lambda x: x["response_length"], reverse=True)
|
||||
for i, result in enumerate(sorted_by_length, 1):
|
||||
print(f"{i}. {result['name']}({result['type']}) - {result['response_length']}字符")
|
||||
else:
|
||||
print("\n没有成功完成测试的模型。")
|
||||
|
||||
if len(test_results) != len(successful_results):
|
||||
print("\n测试失败的模型:")
|
||||
failed_results = [r for r in test_results if r["status"] != "成功" or r["first_response_time"] is None]
|
||||
for result in failed_results:
|
||||
print(f"- {result['name']}({result['type']}): {result['status']}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 运行异步主函数
|
||||
asyncio.run(main())
|
||||
|
||||
+1
-25
@@ -2,8 +2,6 @@ import os
|
||||
import sys
|
||||
from config.logger import setup_logging
|
||||
import importlib
|
||||
from datetime import datetime
|
||||
from core.utils.util import read_config, get_project_dir
|
||||
|
||||
logger = setup_logging()
|
||||
|
||||
@@ -16,26 +14,4 @@ def create_instance(class_name, *args, **kwargs):
|
||||
sys.modules[lib_name] = importlib.import_module(f'{lib_name}')
|
||||
return sys.modules[lib_name].TTSProvider(*args, **kwargs)
|
||||
|
||||
raise ValueError(f"不支持的TTS类型: {class_name},请检查该配置的type是否设置正确")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
响应速度测试
|
||||
"""
|
||||
config = read_config(get_project_dir() + "config.yaml")
|
||||
tts = create_instance(
|
||||
config["selected_module"]["TTS"]
|
||||
if not 'type' in config["TTS"][config["selected_module"]["TTS"]]
|
||||
else
|
||||
config["TTS"][config["selected_module"]["TTS"]]["type"],
|
||||
config["TTS"][config["selected_module"]["TTS"]],
|
||||
config["delete_audio"]
|
||||
)
|
||||
tts.output_file = get_project_dir() + tts.output_file
|
||||
start = datetime.now()
|
||||
file_path = tts.to_tts("你好,测试,我是人工智能小智")
|
||||
print("语音合成耗时:" + str(datetime.now() - start))
|
||||
start = datetime.now()
|
||||
tts.wav_to_opus_data(file_path)
|
||||
print("语音opus耗时:" + str(datetime.now() - start))
|
||||
raise ValueError(f"不支持的TTS类型: {class_name},请检查该配置的type是否设置正确")
|
||||
@@ -34,13 +34,6 @@ def write_json_file(file_path, data):
|
||||
json.dump(data, file, ensure_ascii=False, indent=4)
|
||||
|
||||
|
||||
def is_segment(tokens):
|
||||
if tokens[-1] in (",", ".", "?", ",", "。", "?", "!", "!", ";", ";", ":", ":"):
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
def is_punctuation_or_emoji(char):
|
||||
"""检查字符是否为空格、指定标点或表情符号"""
|
||||
# 定义需要去除的中英文标点(包括全角/半角)
|
||||
|
||||
Reference in New Issue
Block a user