Merge pull request #2291 from xinnan-tech/performance

update:更新测试工具以测试讯飞流式ASR、TTS首词响应时间
This commit is contained in:
hrz
2025-09-26 19:06:28 +08:00
committed by GitHub
2 changed files with 288 additions and 1 deletions
@@ -11,6 +11,12 @@ from tabulate import tabulate
from config.settings import load_config from config.settings import load_config
import tempfile import tempfile
import wave import wave
import hmac
import base64
import hashlib
from datetime import datetime
from wsgiref.handlers import format_date_time
from time import mktime
description = "流式ASR首词延迟测试" description = "流式ASR首词延迟测试"
try: try:
import dashscope import dashscope
@@ -259,6 +265,137 @@ class QwenASRFlashTester(BaseASRTester):
return self._calculate_result("通义千问ASR", latencies, test_count) return self._calculate_result("通义千问ASR", latencies, test_count)
class XunfeiStreamASRTester(BaseASRTester):
def __init__(self):
super().__init__("XunfeiStreamASR")
def _create_url(self):
"""生成讯飞ASR认证URL"""
url = 'ws://iat.cn-huabei-1.xf-yun.com/v1'
# 生成RFC1123格式的时间戳
now = datetime.now()
date = format_date_time(mktime(now.timetuple()))
# 拼接字符串
signature_origin = "host: " + "iat.cn-huabei-1.xf-yun.com" + "\n"
signature_origin += "date: " + date + "\n"
signature_origin += "GET " + "/v1 " + "HTTP/1.1"
# 进行hmac-sha256进行加密
signature_sha = hmac.new(self.asr_config["api_secret"].encode('utf-8'), signature_origin.encode('utf-8'),
digestmod=hashlib.sha256).digest()
signature_sha = base64.b64encode(signature_sha).decode(encoding='utf-8')
authorization_origin = "api_key=\"%s\", algorithm=\"%s\", headers=\"%s\", signature=\"%s\"" % (
self.asr_config["api_key"], "hmac-sha256", "host date request-line", signature_sha)
authorization = base64.b64encode(authorization_origin.encode('utf-8')).decode(encoding='utf-8')
# 将请求的鉴权参数组合为字典
v = {
"authorization": authorization,
"date": date,
"host": "iat.cn-huabei-1.xf-yun.com"
}
# 拼接鉴权参数,生成url
url = url + '?' + parse.urlencode(v)
return url
async def test(self, test_count=5):
if not self.test_audio_files:
return {"name": "讯飞流式ASR", "latency": 0, "status": "失败: 未找到测试音频"}
if not self.asr_config:
return {"name": "讯飞流式ASR", "latency": 0, "status": "失败: 未配置"}
# 检查必要的配置参数
required_keys = ["app_id", "api_key", "api_secret"]
for key in required_keys:
if key not in self.asr_config:
return {"name": "讯飞流式ASR", "latency": 0, "status": f"失败: 缺少配置项 {key}"}
latencies = []
STATUS_FIRST_FRAME = 0
for i in range(test_count):
try:
# 生成认证URL
ws_url = self._create_url()
# 获取音频数据
audio_data = self.test_audio_files[0]['data']
if audio_data.startswith(b'RIFF'):
audio_data = audio_data[44:] # 跳过WAV文件头
# 识别参数
iat_params = {
"domain": self.asr_config.get("domain", "slm"),
"language": self.asr_config.get("language", "zh_cn"),
"accent": self.asr_config.get("accent", "mandarin"),
"dwa": self.asr_config.get("dwa", "wpgs"),
"result": {
"encoding": "utf8",
"compress": "raw",
"format": "plain"
}
}
# 准备首帧数据
first_frame_data = {
"header": {
"status": STATUS_FIRST_FRAME,
"app_id": self.asr_config["app_id"]
},
"parameter": {
"iat": iat_params
},
"payload": {
"audio": {
"audio": base64.b64encode(audio_data[:960]).decode('utf-8'),
"sample_rate": 16000,
"encoding": "raw"
}
}
}
# 启动连接并测量时间
start_time = time.time()
async with websockets.connect(
ws_url,
max_size=1000000000,
ping_interval=None,
ping_timeout=None,
close_timeout=30,
) as ws:
# 发送首帧数据
await ws.send(json.dumps(first_frame_data, ensure_ascii=False))
print(f"[讯飞ASR] 第{i+1}次测试:已发送首帧,等待响应...")
# 直接等待第一个响应并计算延迟
# 参考豆包和通义千问的实现方式,简化逻辑
response_received = False
while not response_received:
try:
# 设置较大的超时时间
response = await asyncio.wait_for(ws.recv(), timeout=30.0)
# 收到响应立即计算延迟,不管内容是什么
# 这样可以准确测量首包到达时间
latency = time.time() - start_time
latencies.append(latency)
response_received = True
print(f"[讯飞ASR] 第{i+1}次测试:收到首包响应,延迟: {latency:.3f}s")
break
except asyncio.TimeoutError:
print(f"[讯飞ASR] 第{i+1}次测试:响应超时")
raise Exception("获取响应超时")
except Exception as e:
print(f"[讯飞ASR] 第{i+1}次测试失败: {str(e)}")
latencies.append(0)
return self._calculate_result("讯飞流式ASR", latencies, test_count)
class ASRPerformanceSuite: class ASRPerformanceSuite:
def __init__(self): def __init__(self):
self.testers = [] self.testers = []
@@ -272,7 +409,8 @@ class ASRPerformanceSuite:
except Exception as e: except Exception as e:
name_map = { name_map = {
"DoubaoStreamASRTester": "豆包流式ASR", "DoubaoStreamASRTester": "豆包流式ASR",
"QwenASRFlashTester": "通义千问ASR" "QwenASRFlashTester": "通义千问ASR",
"XunfeiStreamASRTester": "讯飞流式ASR"
} }
name = name_map.get(tester_class.__name__, tester_class.__name__) name = name_map.get(tester_class.__name__, tester_class.__name__)
print(f"跳过 {name}: {str(e)}") print(f"跳过 {name}: {str(e)}")
@@ -326,6 +464,7 @@ async def main():
suite = ASRPerformanceSuite() suite = ASRPerformanceSuite()
suite.register_tester(DoubaoStreamASRTester) suite.register_tester(DoubaoStreamASRTester)
suite.register_tester(QwenASRFlashTester) suite.register_tester(QwenASRFlashTester)
suite.register_tester(XunfeiStreamASRTester)
await suite.run(args.count) await suite.run(args.count)
@@ -4,6 +4,11 @@ import json
import uuid import uuid
import aiohttp import aiohttp
import websockets import websockets
import hmac
import base64
import hashlib
import asyncio
from urllib.parse import urlparse, urlencode
from tabulate import tabulate from tabulate import tabulate
from config.settings import load_config from config.settings import load_config
@@ -286,6 +291,144 @@ class StreamTTSPerformanceTester:
return self._calculate_result("LinkeraiTTS", latencies, test_count) return self._calculate_result("LinkeraiTTS", latencies, test_count)
async def test_xunfei_tts(self, text=None, test_count=5):
"""测试讯飞流式TTS首词延迟(测试多次取平均)"""
text = text or self.test_texts[0]
latencies = []
for i in range(test_count):
try:
# 修正配置节点名称,与配置文件中的XunFeiTTS匹配
tts_config = self.config["TTS"]["XunFeiTTS"]
app_id = tts_config["app_id"]
api_key = tts_config["api_key"]
api_secret = tts_config["api_secret"]
api_url = tts_config.get("api_url", "wss://cbm01.cn-huabei-1.xf-yun.com/v1/private/mcd9m97e6")
voice = tts_config.get("voice", "x5_lingxiaoxuan_flow")
# 生成认证URL
auth_url = self._create_xunfei_auth_url(api_key, api_secret, api_url)
async with websockets.connect(
auth_url,
ping_interval=30,
ping_timeout=10,
close_timeout=10,
max_size=1000000000
) as ws:
# 构造请求
request = self._build_xunfei_request(app_id, text, voice)
# 发送请求后立即计时,确保准确测量从发送文本到接收首块的时间
await ws.send(json.dumps(request))
start_time = time.time()
# 等待第一个音频数据块
first_audio_received = False
while not first_audio_received:
msg = await asyncio.wait_for(ws.recv(), timeout=10)
data = json.loads(msg)
header = data.get("header", {})
code = header.get("code")
if code != 0:
message = header.get("message", "未知错误")
raise Exception(f"合成失败: {code} - {message}")
payload = data.get("payload", {})
audio_payload = payload.get("audio", {})
if audio_payload:
status = audio_payload.get("status", 0)
audio_data = audio_payload.get("audio", "")
if status == 1 and audio_data:
# 收到第一个音频数据块
latency = time.time() - start_time
latencies.append(latency)
first_audio_received = True
break
except Exception as e:
latencies.append(0)
return self._calculate_result("讯飞TTS", latencies, test_count)
def _create_xunfei_auth_url(self, api_key, api_secret, api_url):
"""生成讯飞WebSocket认证URL"""
parsed_url = urlparse(api_url)
host = parsed_url.netloc
path = parsed_url.path
# 获取UTC时间,讯飞要求使用RFC1123格式
now = time.gmtime()
date = time.strftime('%a, %d %b %Y %H:%M:%S GMT', now)
# 构造签名字符串
signature_origin = f"host: {host}\ndate: {date}\nGET {path} HTTP/1.1"
# 计算签名
signature_sha = hmac.new(
api_secret.encode('utf-8'),
signature_origin.encode('utf-8'),
digestmod=hashlib.sha256
).digest()
signature_sha_base64 = base64.b64encode(signature_sha).decode(encoding='utf-8')
# 构造authorization
authorization_origin = f'api_key="{api_key}", algorithm="hmac-sha256", headers="host date request-line", signature="{signature_sha_base64}"'
authorization = base64.b64encode(authorization_origin.encode('utf-8')).decode(encoding='utf-8')
# 构造最终的WebSocket URL
v = {
"authorization": authorization,
"date": date,
"host": host
}
url = api_url + '?' + urlencode(v)
return url
def _build_xunfei_request(self, app_id, text, voice):
"""构建讯飞TTS请求结构"""
return {
"header": {
"app_id": app_id,
"status": 2,
},
"parameter": {
"oral": {
"oral_level": "mid",
"spark_assist": 1,
"stop_split": 0,
"remain": 0
},
"tts": {
"vcn": voice,
"speed": 50,
"volume": 50,
"pitch": 50,
"bgs": 0,
"reg": 0,
"rdn": 0,
"rhy": 0,
"audio": {
"encoding": "raw",
"sample_rate": 24000,
"channels": 1,
"bit_depth": 16,
"frame_size": 0
}
}
},
"payload": {
"text": {
"encoding": "utf8",
"compress": "raw",
"format": "plain",
"status": 2,
"seq": 1,
"text": base64.b64encode(text.encode('utf-8')).decode('utf-8')
}
}
}
def _calculate_result(self, service_name, latencies, test_count): def _calculate_result(self, service_name, latencies, test_count):
"""计算测试结果""" """计算测试结果"""
@@ -368,6 +511,11 @@ class StreamTTSPerformanceTester:
result = await self.test_indexstream_tts(test_text, test_count) result = await self.test_indexstream_tts(test_text, test_count)
self.results.append(result) self.results.append(result)
# 测试讯飞TTS
if self.config.get("TTS", {}).get("XunFeiTTS"):
result = await self.test_xunfei_tts(test_text, test_count)
self.results.append(result)
# 打印结果 # 打印结果
self._print_results(test_text, test_count) self._print_results(test_text, test_count)