移除TTS相关代码
parent
5b4ab416e2
commit
86c9f1d415
|
|
@ -1,6 +1,6 @@
|
||||||
# Qwen3-ASR 并发性能测试脚本
|
# Qwen3-ASR 并发性能测试脚本
|
||||||
|
|
||||||
测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。
|
测试 ASR WebSocket 服务在不同并发级别下的性能表现。
|
||||||
|
|
||||||
## 依赖
|
## 依赖
|
||||||
|
|
||||||
|
|
@ -21,15 +21,8 @@ python start.py
|
||||||
### 2. 运行测试
|
### 2. 运行测试
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# 完整测试 (ASR + TTS)
|
|
||||||
.venv/bin/python -m scripts.benchmark.run --audio-file /path/to/audio.wav
|
.venv/bin/python -m scripts.benchmark.run --audio-file /path/to/audio.wav
|
||||||
|
|
||||||
# 仅测试 TTS (无需音频文件)
|
|
||||||
.venv/bin/python -m scripts.benchmark.run --test-type tts
|
|
||||||
|
|
||||||
# 仅测试 ASR
|
|
||||||
.venv/bin/python -m scripts.benchmark.run --audio-file /path/to/audio.wav --test-type asr
|
|
||||||
|
|
||||||
# Qwen Rust CPU 固定配置跑测(固定 VAD 分段)
|
# Qwen Rust CPU 固定配置跑测(固定 VAD 分段)
|
||||||
.venv/bin/python -m scripts.benchmark.qwen_rust_sensitivity \
|
.venv/bin/python -m scripts.benchmark.qwen_rust_sensitivity \
|
||||||
--audio-file /path/to/audio.wav
|
--audio-file /path/to/audio.wav
|
||||||
|
|
@ -41,12 +34,10 @@ python start.py
|
||||||
|------|--------|------|
|
|------|--------|------|
|
||||||
| `--host` | localhost | 服务器主机名 |
|
| `--host` | localhost | 服务器主机名 |
|
||||||
| `--port` | 8000 | 服务器端口 |
|
| `--port` | 8000 | 服务器端口 |
|
||||||
| `--audio-file` | - | ASR 测试音频文件路径 (测试 ASR 时必需) |
|
| `--audio-file` | - | ASR 测试音频文件路径(必需) |
|
||||||
| `--test-type` | both | 测试类型: `asr` / `tts` / `both` |
|
|
||||||
| `--concurrency` | 5 10 20 50 | 并发级别列表 |
|
| `--concurrency` | 5 10 20 50 | 并发级别列表 |
|
||||||
| `--output` | ./benchmark_results | 报告输出目录 |
|
| `--output` | ./benchmark_results | 报告输出目录 |
|
||||||
| `--timeout` | 120 | 请求超时时间 (秒) |
|
| `--timeout` | 120 | 请求超时时间 (秒) |
|
||||||
| `--voice` | 中文女 | TTS 测试音色 |
|
|
||||||
|
|
||||||
## Qwen Rust CPU 固定配置跑测
|
## Qwen Rust CPU 固定配置跑测
|
||||||
|
|
||||||
|
|
@ -89,14 +80,6 @@ QWEN_RUST_ALIGN_CONCURRENCY=4 \
|
||||||
--markdown-out temp/qwen_rust_runtime_config_report.md
|
--markdown-out temp/qwen_rust_runtime_config_report.md
|
||||||
```
|
```
|
||||||
|
|
||||||
### TTS 流式模拟配置 (config.py)
|
|
||||||
|
|
||||||
TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、顿号等)分割文本逐步发送:
|
|
||||||
|
|
||||||
| 配置项 | 默认值 | 说明 |
|
|
||||||
|--------|--------|------|
|
|
||||||
| `tts_chunk_interval` | 0.05 | 发送间隔秒数 (模拟 LLM 生成速度) |
|
|
||||||
|
|
||||||
## 使用示例
|
## 使用示例
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -105,16 +88,11 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
||||||
--audio-file test.wav \
|
--audio-file test.wav \
|
||||||
--concurrency 5 10 20 50 100
|
--concurrency 5 10 20 50 100
|
||||||
|
|
||||||
# 连接远程服务器
|
# 连接远程 ASR 服务
|
||||||
.venv/bin/python -m scripts.benchmark.run \
|
.venv/bin/python -m scripts.benchmark.run \
|
||||||
--host 192.168.1.100 \
|
--host 192.168.1.100 \
|
||||||
--port 8000 \
|
--port 8000 \
|
||||||
--test-type tts
|
--audio-file test.wav
|
||||||
|
|
||||||
# 使用不同音色测试 TTS
|
|
||||||
.venv/bin/python -m scripts.benchmark.run \
|
|
||||||
--test-type tts \
|
|
||||||
--voice 中文男
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## 测试指标
|
## 测试指标
|
||||||
|
|
@ -124,11 +102,6 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
||||||
- **总处理时间**: 从开始到识别完成的总时间
|
- **总处理时间**: 从开始到识别完成的总时间
|
||||||
- **RTF**: 处理时间 / 音频时长 (小于 1.0 表示快于实时)
|
- **RTF**: 处理时间 / 音频时长 (小于 1.0 表示快于实时)
|
||||||
|
|
||||||
### TTS 指标
|
|
||||||
- **首包延迟**: 从发送文本到收到第一个音频块的时间
|
|
||||||
- **总合成时间**: 从开始到合成完成的总时间
|
|
||||||
- **RTF**: 合成时间 / 生成音频时长
|
|
||||||
|
|
||||||
### 统计维度
|
### 统计维度
|
||||||
每个指标计算: 平均值 (Avg)、P50、P95、P99、最大值 (Max)
|
每个指标计算: 平均值 (Avg)、P50、P95、P99、最大值 (Max)
|
||||||
|
|
||||||
|
|
@ -139,8 +112,8 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
||||||
```
|
```
|
||||||
benchmark_results/
|
benchmark_results/
|
||||||
├── benchmark_report_20241202_143000.md # Markdown 报告
|
├── benchmark_report_20241202_143000.md # Markdown 报告
|
||||||
├── first_latency_20241202_143000.png # 首次响应延迟图
|
├── first_latency_20241202_143000.png # ASR 首次响应延迟图
|
||||||
├── rtf_20241202_143000.png # RTF 对比图
|
├── rtf_20241202_143000.png # ASR RTF 图
|
||||||
├── throughput_20241202_143000.png # 吞吐量图
|
├── throughput_20241202_143000.png # 吞吐量图
|
||||||
└── total_time_20241202_143000.png # 总时间图
|
└── total_time_20241202_143000.png # 总时间图
|
||||||
```
|
```
|
||||||
|
|
@ -174,7 +147,6 @@ scripts/benchmark/
|
||||||
├── clients/
|
├── clients/
|
||||||
│ ├── base_client.py # WebSocket 客户端基类
|
│ ├── base_client.py # WebSocket 客户端基类
|
||||||
│ ├── asr_client.py # ASR 测试客户端
|
│ ├── asr_client.py # ASR 测试客户端
|
||||||
│ └── tts_client.py # TTS 测试客户端
|
|
||||||
├── metrics/
|
├── metrics/
|
||||||
│ ├── models.py # 指标数据类
|
│ ├── models.py # 指标数据类
|
||||||
│ └── statistics.py # 统计计算
|
│ └── statistics.py # 统计计算
|
||||||
|
|
@ -182,17 +154,14 @@ scripts/benchmark/
|
||||||
│ ├── markdown_reporter.py # Markdown 报告生成
|
│ ├── markdown_reporter.py # Markdown 报告生成
|
||||||
│ └── chart_generator.py # 图表生成
|
│ └── chart_generator.py # 图表生成
|
||||||
└── utils/
|
└── utils/
|
||||||
├── audio_utils.py # 音频文件处理
|
└── audio_utils.py # 音频文件处理
|
||||||
└── text_generator.py # 测试文本生成
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## 注意事项
|
## 注意事项
|
||||||
|
|
||||||
1. **ASR 测试需要音频文件**: 建议使用 1 分钟左右的音频,格式支持 wav/mp3 等常见格式
|
1. **ASR 测试需要音频文件**: 建议使用 1 分钟左右的音频,格式支持 wav/mp3 等常见格式
|
||||||
2. **TTS 测试自动生成文本**: 使用内置的中文随机句子生成器,无需额外准备
|
2. **并发测试会占用资源**: 高并发测试时请确保服务器有足够资源
|
||||||
3. **TTS 模拟流式输入**: 测试会按标点符号(逗号、句号、顿号等)分割文本逐步发送,模拟 LLM 流式输出场景
|
3. **RTF 解读**:
|
||||||
4. **并发测试会占用资源**: 高并发测试时请确保服务器有足够资源
|
|
||||||
5. **RTF 解读**:
|
|
||||||
- RTF < 1.0: 处理速度快于实时,性能良好
|
- RTF < 1.0: 处理速度快于实时,性能良好
|
||||||
- RTF ≈ 1.0: 刚好实时处理
|
- RTF ≈ 1.0: 刚好实时处理
|
||||||
- RTF > 1.0: 处理速度慢于实时,可能出现延迟累积
|
- RTF > 1.0: 处理速度慢于实时,可能出现延迟累积
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@
|
||||||
"""
|
"""
|
||||||
Qwen3-ASR 并发性能测试脚本
|
Qwen3-ASR 并发性能测试脚本
|
||||||
|
|
||||||
用于测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。
|
用于测试 ASR WebSocket 服务在不同并发级别下的性能表现。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
__version__ = "1.0.0"
|
__version__ = "1.0.0"
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,5 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
from .base_client import BaseWebSocketClient
|
from .base_client import BaseWebSocketClient
|
||||||
from .asr_client import ASRWebSocketClient
|
from .asr_client import ASRWebSocketClient
|
||||||
from .tts_client import TTSWebSocketClient
|
|
||||||
|
|
||||||
__all__ = ["BaseWebSocketClient", "ASRWebSocketClient", "TTSWebSocketClient"]
|
__all__ = ["BaseWebSocketClient", "ASRWebSocketClient"]
|
||||||
|
|
|
||||||
|
|
@ -1,251 +0,0 @@
|
||||||
# -*- coding: utf-8 -*-
|
|
||||||
"""
|
|
||||||
TTS WebSocket 测试客户端
|
|
||||||
|
|
||||||
参考 realtime-llm/backend/cti_websocket_handler.py 中的调用方式。
|
|
||||||
模拟 LLM 流式输出场景,按句子发送文本进行合成。
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import json
|
|
||||||
import time
|
|
||||||
import logging
|
|
||||||
import wave
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from .base_client import BaseWebSocketClient
|
|
||||||
from ..metrics.models import TTSMetrics
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
# 协议常量
|
|
||||||
TTS_NAMESPACE = "FlowingSpeechSynthesizer"
|
|
||||||
MSG_START_SYNTHESIS = "StartSynthesis"
|
|
||||||
MSG_RUN_SYNTHESIS = "RunSynthesis"
|
|
||||||
MSG_STOP_SYNTHESIS = "StopSynthesis"
|
|
||||||
MSG_SYNTHESIS_STARTED = "SynthesisStarted"
|
|
||||||
MSG_SENTENCE_BEGIN = "SentenceBegin"
|
|
||||||
MSG_SENTENCE_END = "SentenceEnd"
|
|
||||||
MSG_SYNTHESIS_COMPLETED = "SynthesisCompleted"
|
|
||||||
MSG_TASK_FAILED = "TaskFailed"
|
|
||||||
|
|
||||||
|
|
||||||
class TTSWebSocketClient(BaseWebSocketClient):
|
|
||||||
"""TTS WebSocket 测试客户端 (模拟流式文本输入)"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
ws_url: str,
|
|
||||||
text: str,
|
|
||||||
voice: str = "中文女",
|
|
||||||
audio_format: str = "PCM",
|
|
||||||
sample_rate: int = 22050,
|
|
||||||
timeout: float = 120.0,
|
|
||||||
chunk_interval: float = 0.05, # 发送间隔 (秒),模拟 LLM 生成速度
|
|
||||||
debug: bool = False, # 调试模式
|
|
||||||
save_audio_dir: Optional[Path] = None, # 保存音频的目录
|
|
||||||
):
|
|
||||||
super().__init__(ws_url, timeout)
|
|
||||||
self.text = text
|
|
||||||
self.voice = voice
|
|
||||||
self.audio_format = audio_format
|
|
||||||
self.sample_rate = sample_rate
|
|
||||||
self.chunk_interval = chunk_interval
|
|
||||||
self.debug = debug
|
|
||||||
self.save_audio_dir = save_audio_dir
|
|
||||||
self._audio_chunks = [] # 存储接收到的音频数据
|
|
||||||
|
|
||||||
def _log(self, msg: str):
|
|
||||||
"""调试日志"""
|
|
||||||
if self.debug:
|
|
||||||
logger.info(f"[{self.task_id[:8]}] {msg}")
|
|
||||||
|
|
||||||
async def run_test(self) -> TTSMetrics:
|
|
||||||
"""执行 TTS 测试"""
|
|
||||||
metrics = TTSMetrics(
|
|
||||||
request_id=self.task_id,
|
|
||||||
concurrency_level=0,
|
|
||||||
start_time=time.perf_counter(),
|
|
||||||
text_length=len(self.text),
|
|
||||||
sample_rate=self.sample_rate,
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
await asyncio.wait_for(
|
|
||||||
self._run_tts_session(metrics),
|
|
||||||
timeout=self.timeout,
|
|
||||||
)
|
|
||||||
metrics.success = True
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
metrics.error_message = "Timeout"
|
|
||||||
logger.warning(f"TTS 请求超时: {self.task_id}")
|
|
||||||
except Exception as e:
|
|
||||||
metrics.error_message = str(e)
|
|
||||||
logger.warning(f"TTS 请求失败: {self.task_id}, 错误: {e}")
|
|
||||||
finally:
|
|
||||||
await self.close()
|
|
||||||
|
|
||||||
return metrics
|
|
||||||
|
|
||||||
async def _run_tts_session(self, metrics: TTSMetrics) -> None:
|
|
||||||
"""运行完整的 TTS 会话"""
|
|
||||||
self._log(f"连接 {self.ws_url}")
|
|
||||||
await self.connect()
|
|
||||||
|
|
||||||
# 用于同步的事件
|
|
||||||
started_event = asyncio.Event()
|
|
||||||
completed_event = asyncio.Event()
|
|
||||||
error_message = None
|
|
||||||
|
|
||||||
# 1. 发送 StartSynthesis
|
|
||||||
await self._send_start_synthesis()
|
|
||||||
|
|
||||||
# 2. 启动接收任务
|
|
||||||
async def receive_loop():
|
|
||||||
nonlocal error_message
|
|
||||||
while True:
|
|
||||||
try:
|
|
||||||
response = await self.receive()
|
|
||||||
except Exception as e:
|
|
||||||
self._log(f"接收异常: {e}")
|
|
||||||
break
|
|
||||||
|
|
||||||
if isinstance(response, bytes):
|
|
||||||
if metrics.first_chunk_time is None:
|
|
||||||
metrics.first_chunk_time = time.perf_counter()
|
|
||||||
metrics.audio_bytes_received += len(response)
|
|
||||||
# 收集音频数据用于保存
|
|
||||||
if self.save_audio_dir:
|
|
||||||
self._audio_chunks.append(response)
|
|
||||||
self._log(f"← 收到音频: {len(response)} bytes")
|
|
||||||
|
|
||||||
elif isinstance(response, str):
|
|
||||||
try:
|
|
||||||
data = json.loads(response)
|
|
||||||
header = data.get("header", {})
|
|
||||||
name = header.get("name", "")
|
|
||||||
status = header.get("status", 0)
|
|
||||||
|
|
||||||
self._log(f"← 收到事件: {name} (status={status})")
|
|
||||||
|
|
||||||
if name == MSG_SYNTHESIS_STARTED:
|
|
||||||
started_event.set()
|
|
||||||
|
|
||||||
elif name == MSG_SENTENCE_END:
|
|
||||||
metrics.sentence_end_time = time.perf_counter()
|
|
||||||
|
|
||||||
elif name == MSG_SYNTHESIS_COMPLETED:
|
|
||||||
completed_event.set()
|
|
||||||
break
|
|
||||||
|
|
||||||
elif name == MSG_TASK_FAILED:
|
|
||||||
status_text = header.get("status_text", "Unknown error")
|
|
||||||
error_message = f"TaskFailed: {status_text}"
|
|
||||||
self._log(f"← 错误: {status_text}")
|
|
||||||
completed_event.set()
|
|
||||||
break
|
|
||||||
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
receive_task = asyncio.create_task(receive_loop())
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 3. 等待 SynthesisStarted
|
|
||||||
self._log("等待 SynthesisStarted...")
|
|
||||||
await asyncio.wait_for(started_event.wait(), timeout=10.0)
|
|
||||||
self._log("收到 SynthesisStarted")
|
|
||||||
|
|
||||||
# 4. 发送文本 - 直接发送完整文本,不分割
|
|
||||||
# 参考 CTI 客户端:发送完整句子而不是切分片段
|
|
||||||
await self._send_run_synthesis(self.text)
|
|
||||||
|
|
||||||
# 5. 发送 StopSynthesis
|
|
||||||
await self._send_stop_synthesis()
|
|
||||||
|
|
||||||
# 6. 等待 SynthesisCompleted
|
|
||||||
self._log("等待 SynthesisCompleted...")
|
|
||||||
await completed_event.wait()
|
|
||||||
|
|
||||||
if error_message:
|
|
||||||
raise Exception(error_message)
|
|
||||||
|
|
||||||
metrics.complete_time = time.perf_counter()
|
|
||||||
self._log(f"完成! 收到 {metrics.audio_bytes_received} bytes 音频")
|
|
||||||
|
|
||||||
# 保存音频文件
|
|
||||||
if self.save_audio_dir and self._audio_chunks:
|
|
||||||
self._save_audio()
|
|
||||||
|
|
||||||
finally:
|
|
||||||
if not receive_task.done():
|
|
||||||
receive_task.cancel()
|
|
||||||
try:
|
|
||||||
await receive_task
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def _send_start_synthesis(self) -> None:
|
|
||||||
"""发送 StartSynthesis 消息"""
|
|
||||||
message = {
|
|
||||||
"header": self._create_header(MSG_START_SYNTHESIS, TTS_NAMESPACE),
|
|
||||||
"payload": {
|
|
||||||
"voice": self.voice,
|
|
||||||
"format": self.audio_format,
|
|
||||||
"sample_rate": self.sample_rate,
|
|
||||||
"volume": 50,
|
|
||||||
"speech_rate": 0,
|
|
||||||
"pitch_rate": 0,
|
|
||||||
"platform": "python",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
self._log(f"→ 发送 StartSynthesis (voice={self.voice}, format={self.audio_format})")
|
|
||||||
await self.send_json(message)
|
|
||||||
|
|
||||||
async def _send_run_synthesis(self, text: str) -> None:
|
|
||||||
"""发送 RunSynthesis 消息"""
|
|
||||||
message = {
|
|
||||||
"header": self._create_header(MSG_RUN_SYNTHESIS, TTS_NAMESPACE),
|
|
||||||
"payload": {
|
|
||||||
"text": text,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
# 截断显示
|
|
||||||
display_text = text[:50] + "..." if len(text) > 50 else text
|
|
||||||
self._log(f"→ 发送 RunSynthesis: \"{display_text}\" ({len(text)} chars)")
|
|
||||||
await self.send_json(message)
|
|
||||||
|
|
||||||
async def _send_stop_synthesis(self) -> None:
|
|
||||||
"""发送 StopSynthesis 消息"""
|
|
||||||
message = {
|
|
||||||
"header": self._create_header(MSG_STOP_SYNTHESIS, TTS_NAMESPACE),
|
|
||||||
}
|
|
||||||
self._log("→ 发送 StopSynthesis")
|
|
||||||
await self.send_json(message)
|
|
||||||
|
|
||||||
def _save_audio(self) -> None:
|
|
||||||
"""保存收到的音频数据为 WAV 文件"""
|
|
||||||
if self.save_audio_dir is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 合并所有音频块
|
|
||||||
audio_data = b"".join(self._audio_chunks)
|
|
||||||
if not audio_data:
|
|
||||||
return
|
|
||||||
|
|
||||||
# 生成文件名
|
|
||||||
filename = f"{self.task_id[:8]}_{len(self.text)}chars.wav"
|
|
||||||
filepath = self.save_audio_dir / filename
|
|
||||||
|
|
||||||
# PCM 数据保存为 WAV
|
|
||||||
with wave.open(str(filepath), 'wb') as wav_file:
|
|
||||||
wav_file.setnchannels(1) # 单声道
|
|
||||||
wav_file.setsampwidth(2) # 16位 = 2字节
|
|
||||||
wav_file.setframerate(self.sample_rate)
|
|
||||||
wav_file.writeframes(audio_data)
|
|
||||||
|
|
||||||
self._log(f"音频已保存: {filepath}")
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"保存音频失败: {e}")
|
|
||||||
|
|
@ -15,7 +15,7 @@ class TestConfig:
|
||||||
# 服务器配置
|
# 服务器配置
|
||||||
host: str = "localhost"
|
host: str = "localhost"
|
||||||
port: int = 8000
|
port: int = 8000
|
||||||
timeout_seconds: float = 300.0 # 默认 5 分钟,并发 TTS 可能需要更长时间
|
timeout_seconds: float = 300.0
|
||||||
warmup_requests: int = 3
|
warmup_requests: int = 3
|
||||||
|
|
||||||
# 并发配置
|
# 并发配置
|
||||||
|
|
@ -27,14 +27,6 @@ class TestConfig:
|
||||||
asr_chunk_size: int = 9600 # 600ms @ 16kHz
|
asr_chunk_size: int = 9600 # 600ms @ 16kHz
|
||||||
asr_format: str = "pcm"
|
asr_format: str = "pcm"
|
||||||
|
|
||||||
# TTS 配置
|
|
||||||
tts_text_count: int = 50 # 预生成的测试文本数量
|
|
||||||
tts_text_length_range: tuple = (50, 100) # 文本字符数范围
|
|
||||||
tts_voice: str = "中文女"
|
|
||||||
tts_format: str = "PCM"
|
|
||||||
tts_sample_rate: int = 22050
|
|
||||||
tts_chunk_interval: float = 0.05 # 发送间隔秒数 (模拟 LLM 生成速度)
|
|
||||||
|
|
||||||
# 输出配置
|
# 输出配置
|
||||||
output_dir: Path = field(default_factory=lambda: Path("./benchmark_results"))
|
output_dir: Path = field(default_factory=lambda: Path("./benchmark_results"))
|
||||||
report_name: str = "benchmark_report"
|
report_name: str = "benchmark_report"
|
||||||
|
|
@ -49,22 +41,13 @@ class TestConfig:
|
||||||
"""ASR WebSocket URL"""
|
"""ASR WebSocket URL"""
|
||||||
return f"{self.ws_base_url}/ws/v1/asr"
|
return f"{self.ws_base_url}/ws/v1/asr"
|
||||||
|
|
||||||
@property
|
def validate(self) -> None:
|
||||||
def tts_ws_url(self) -> str:
|
|
||||||
"""TTS WebSocket URL"""
|
|
||||||
return f"{self.ws_base_url}/ws/v1/tts"
|
|
||||||
|
|
||||||
def validate(self, test_type: str = "both") -> None:
|
|
||||||
"""
|
"""
|
||||||
验证配置
|
验证配置
|
||||||
|
|
||||||
Args:
|
|
||||||
test_type: 测试类型 (asr/tts/both)
|
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
ValueError: 配置无效
|
ValueError: 配置无效
|
||||||
"""
|
"""
|
||||||
if test_type in ("asr", "both"):
|
|
||||||
if self.asr_audio_file is None:
|
if self.asr_audio_file is None:
|
||||||
raise ValueError("ASR 测试需要提供音频文件路径 (--audio-file)")
|
raise ValueError("ASR 测试需要提供音频文件路径 (--audio-file)")
|
||||||
if not self.asr_audio_file.exists():
|
if not self.asr_audio_file.exists():
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,9 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
from .models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
from .models import ASRMetrics, AggregatedMetrics
|
||||||
from .statistics import calculate_statistics, calculate_percentile
|
from .statistics import calculate_statistics, calculate_percentile
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ASRMetrics",
|
"ASRMetrics",
|
||||||
"TTSMetrics",
|
|
||||||
"AggregatedMetrics",
|
"AggregatedMetrics",
|
||||||
"calculate_statistics",
|
"calculate_statistics",
|
||||||
"calculate_percentile",
|
"calculate_percentile",
|
||||||
|
|
|
||||||
|
|
@ -49,64 +49,10 @@ class ASRMetrics:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class TTSMetrics:
|
|
||||||
"""TTS 单次请求指标"""
|
|
||||||
|
|
||||||
request_id: str
|
|
||||||
concurrency_level: int
|
|
||||||
start_time: float # time.perf_counter()
|
|
||||||
text_length: int = 0
|
|
||||||
sample_rate: int = 22050
|
|
||||||
|
|
||||||
# 时间戳
|
|
||||||
first_chunk_time: Optional[float] = None # 第一个音频二进制块
|
|
||||||
sentence_end_time: Optional[float] = None # SentenceEnd
|
|
||||||
complete_time: Optional[float] = None # SynthesisCompleted
|
|
||||||
|
|
||||||
# 结果
|
|
||||||
audio_bytes_received: int = 0
|
|
||||||
success: bool = False
|
|
||||||
error_message: str = ""
|
|
||||||
|
|
||||||
@property
|
|
||||||
def first_chunk_latency_ms(self) -> Optional[float]:
|
|
||||||
"""首包延迟 (ms)"""
|
|
||||||
if self.first_chunk_time is not None:
|
|
||||||
return (self.first_chunk_time - self.start_time) * 1000
|
|
||||||
return None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def total_synthesis_time_ms(self) -> Optional[float]:
|
|
||||||
"""总合成时间 (ms)"""
|
|
||||||
if self.complete_time is not None:
|
|
||||||
return (self.complete_time - self.start_time) * 1000
|
|
||||||
return None
|
|
||||||
|
|
||||||
@property
|
|
||||||
def estimated_audio_duration_ms(self) -> float:
|
|
||||||
"""估算的音频时长 (基于采样率和字节数)"""
|
|
||||||
if self.audio_bytes_received > 0:
|
|
||||||
# PCM 16-bit mono: 2 bytes per sample
|
|
||||||
samples = self.audio_bytes_received / 2
|
|
||||||
return (samples / self.sample_rate) * 1000
|
|
||||||
return 0.0
|
|
||||||
|
|
||||||
@property
|
|
||||||
def rtf(self) -> Optional[float]:
|
|
||||||
"""RTF = 合成时间 / 生成音频时长"""
|
|
||||||
total_time = self.total_synthesis_time_ms
|
|
||||||
audio_duration = self.estimated_audio_duration_ms
|
|
||||||
if total_time is not None and audio_duration > 0:
|
|
||||||
return total_time / audio_duration
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AggregatedMetrics:
|
class AggregatedMetrics:
|
||||||
"""聚合后的指标 (针对一个并发级别)"""
|
"""聚合后的指标 (针对一个并发级别)"""
|
||||||
|
|
||||||
test_type: str # "asr" or "tts"
|
|
||||||
concurrency_level: int
|
concurrency_level: int
|
||||||
total_requests: int
|
total_requests: int
|
||||||
successful_requests: int
|
successful_requests: int
|
||||||
|
|
|
||||||
|
|
@ -3,10 +3,10 @@
|
||||||
统计计算模块
|
统计计算模块
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import List, Union
|
from typing import List
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from .models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
from .models import ASRMetrics, AggregatedMetrics
|
||||||
|
|
||||||
|
|
||||||
def calculate_percentile(values: List[float], percentile: float) -> float:
|
def calculate_percentile(values: List[float], percentile: float) -> float:
|
||||||
|
|
@ -25,7 +25,7 @@ def calculate_percentile(values: List[float], percentile: float) -> float:
|
||||||
return float(np.percentile(values, percentile))
|
return float(np.percentile(values, percentile))
|
||||||
|
|
||||||
|
|
||||||
def calculate_asr_statistics(
|
def calculate_statistics(
|
||||||
metrics_list: List[ASRMetrics],
|
metrics_list: List[ASRMetrics],
|
||||||
concurrency_level: int,
|
concurrency_level: int,
|
||||||
total_test_time: float,
|
total_test_time: float,
|
||||||
|
|
@ -41,6 +41,9 @@ def calculate_asr_statistics(
|
||||||
Returns:
|
Returns:
|
||||||
聚合后的指标
|
聚合后的指标
|
||||||
"""
|
"""
|
||||||
|
if not metrics_list:
|
||||||
|
raise ValueError("指标列表不能为空")
|
||||||
|
|
||||||
successful = [m for m in metrics_list if m.success]
|
successful = [m for m in metrics_list if m.success]
|
||||||
failed = [m for m in metrics_list if not m.success]
|
failed = [m for m in metrics_list if not m.success]
|
||||||
|
|
||||||
|
|
@ -54,7 +57,6 @@ def calculate_asr_statistics(
|
||||||
rtfs = [m.rtf for m in successful if m.rtf is not None]
|
rtfs = [m.rtf for m in successful if m.rtf is not None]
|
||||||
|
|
||||||
return AggregatedMetrics(
|
return AggregatedMetrics(
|
||||||
test_type="asr",
|
|
||||||
concurrency_level=concurrency_level,
|
concurrency_level=concurrency_level,
|
||||||
total_requests=len(metrics_list),
|
total_requests=len(metrics_list),
|
||||||
successful_requests=len(successful),
|
successful_requests=len(successful),
|
||||||
|
|
@ -79,90 +81,3 @@ def calculate_asr_statistics(
|
||||||
rtf_p99=calculate_percentile(rtfs, 99),
|
rtf_p99=calculate_percentile(rtfs, 99),
|
||||||
rtf_max=max(rtfs) if rtfs else 0.0,
|
rtf_max=max(rtfs) if rtfs else 0.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def calculate_tts_statistics(
|
|
||||||
metrics_list: List[TTSMetrics],
|
|
||||||
concurrency_level: int,
|
|
||||||
total_test_time: float,
|
|
||||||
) -> AggregatedMetrics:
|
|
||||||
"""
|
|
||||||
计算 TTS 指标统计
|
|
||||||
|
|
||||||
Args:
|
|
||||||
metrics_list: TTS 指标列表
|
|
||||||
concurrency_level: 并发级别
|
|
||||||
total_test_time: 总测试时间 (秒)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
聚合后的指标
|
|
||||||
"""
|
|
||||||
successful = [m for m in metrics_list if m.success]
|
|
||||||
failed = [m for m in metrics_list if not m.success]
|
|
||||||
|
|
||||||
# 提取各项指标值
|
|
||||||
first_latencies = [
|
|
||||||
m.first_chunk_latency_ms for m in successful if m.first_chunk_latency_ms is not None
|
|
||||||
]
|
|
||||||
total_times = [
|
|
||||||
m.total_synthesis_time_ms for m in successful if m.total_synthesis_time_ms is not None
|
|
||||||
]
|
|
||||||
rtfs = [m.rtf for m in successful if m.rtf is not None]
|
|
||||||
|
|
||||||
return AggregatedMetrics(
|
|
||||||
test_type="tts",
|
|
||||||
concurrency_level=concurrency_level,
|
|
||||||
total_requests=len(metrics_list),
|
|
||||||
successful_requests=len(successful),
|
|
||||||
failed_requests=len(failed),
|
|
||||||
total_test_time_seconds=total_test_time,
|
|
||||||
# 首包延迟
|
|
||||||
first_latency_avg=float(np.mean(first_latencies)) if first_latencies else 0.0,
|
|
||||||
first_latency_p50=calculate_percentile(first_latencies, 50),
|
|
||||||
first_latency_p95=calculate_percentile(first_latencies, 95),
|
|
||||||
first_latency_p99=calculate_percentile(first_latencies, 99),
|
|
||||||
first_latency_max=max(first_latencies) if first_latencies else 0.0,
|
|
||||||
# 总时间
|
|
||||||
total_time_avg=float(np.mean(total_times)) if total_times else 0.0,
|
|
||||||
total_time_p50=calculate_percentile(total_times, 50),
|
|
||||||
total_time_p95=calculate_percentile(total_times, 95),
|
|
||||||
total_time_p99=calculate_percentile(total_times, 99),
|
|
||||||
total_time_max=max(total_times) if total_times else 0.0,
|
|
||||||
# RTF
|
|
||||||
rtf_avg=float(np.mean(rtfs)) if rtfs else 0.0,
|
|
||||||
rtf_p50=calculate_percentile(rtfs, 50),
|
|
||||||
rtf_p95=calculate_percentile(rtfs, 95),
|
|
||||||
rtf_p99=calculate_percentile(rtfs, 99),
|
|
||||||
rtf_max=max(rtfs) if rtfs else 0.0,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def calculate_statistics(
|
|
||||||
metrics_list: Union[List[ASRMetrics], List[TTSMetrics]],
|
|
||||||
concurrency_level: int,
|
|
||||||
total_test_time: float,
|
|
||||||
) -> AggregatedMetrics:
|
|
||||||
"""
|
|
||||||
通用统计计算函数
|
|
||||||
|
|
||||||
Args:
|
|
||||||
metrics_list: 指标列表 (ASR 或 TTS)
|
|
||||||
concurrency_level: 并发级别
|
|
||||||
total_test_time: 总测试时间 (秒)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
聚合后的指标
|
|
||||||
"""
|
|
||||||
if not metrics_list:
|
|
||||||
raise ValueError("指标列表不能为空")
|
|
||||||
|
|
||||||
if isinstance(metrics_list[0], ASRMetrics):
|
|
||||||
# 类型缩窄:确保类型检查器知道这是 List[ASRMetrics]
|
|
||||||
asr_metrics_list: List[ASRMetrics] = [m for m in metrics_list if isinstance(m, ASRMetrics)]
|
|
||||||
return calculate_asr_statistics(asr_metrics_list, concurrency_level, total_test_time)
|
|
||||||
elif isinstance(metrics_list[0], TTSMetrics):
|
|
||||||
# 类型缩窄:确保类型检查器知道这是 List[TTSMetrics]
|
|
||||||
tts_metrics_list: List[TTSMetrics] = [m for m in metrics_list if isinstance(m, TTSMetrics)]
|
|
||||||
return calculate_tts_statistics(tts_metrics_list, concurrency_level, total_test_time)
|
|
||||||
else:
|
|
||||||
raise TypeError(f"不支持的指标类型: {type(metrics_list[0])}")
|
|
||||||
|
|
|
||||||
|
|
@ -1,250 +1,121 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""
|
"""Generate ASR benchmark charts."""
|
||||||
Matplotlib 图表生成器
|
|
||||||
"""
|
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
import matplotlib
|
import matplotlib
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from ..metrics.models import AggregatedMetrics
|
from ..metrics.models import AggregatedMetrics
|
||||||
|
|
||||||
# 设置中文字体支持
|
# Keep generated labels readable on common Windows and Linux installations.
|
||||||
matplotlib.rcParams['font.sans-serif'] = ['Arial Unicode MS', 'SimHei', 'DejaVu Sans']
|
matplotlib.rcParams["font.sans-serif"] = ["Arial Unicode MS", "SimHei", "DejaVu Sans"]
|
||||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
matplotlib.rcParams["axes.unicode_minus"] = False
|
||||||
|
|
||||||
|
|
||||||
class ChartGenerator:
|
class ChartGenerator:
|
||||||
"""图表生成器"""
|
"""Create the latency, RTF, throughput, and duration ASR charts."""
|
||||||
|
|
||||||
def __init__(self):
|
_COLOR = "#4CAF50"
|
||||||
self.colors = {
|
|
||||||
"asr": "#4CAF50", # 绿色
|
|
||||||
"tts": "#2196F3", # 蓝色
|
|
||||||
}
|
|
||||||
|
|
||||||
def generate_all_charts(
|
def generate_all_charts(
|
||||||
self,
|
self,
|
||||||
asr_results: List[AggregatedMetrics],
|
results: List[AggregatedMetrics],
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
output_dir: Path,
|
output_dir: Path,
|
||||||
timestamp: str,
|
timestamp: str,
|
||||||
) -> List[Path]:
|
) -> List[Path]:
|
||||||
"""
|
|
||||||
生成所有图表
|
|
||||||
|
|
||||||
Args:
|
|
||||||
asr_results: ASR 测试结果
|
|
||||||
tts_results: TTS 测试结果
|
|
||||||
output_dir: 输出目录
|
|
||||||
timestamp: 时间戳
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
生成的图表文件路径列表
|
|
||||||
"""
|
|
||||||
output_dir.mkdir(parents=True, exist_ok=True)
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
generated_files = []
|
charts = (
|
||||||
|
("first_latency", self._generate_latency_chart),
|
||||||
|
("rtf", self._generate_rtf_chart),
|
||||||
|
("throughput", self._generate_throughput_chart),
|
||||||
|
("total_time", self._generate_total_time_chart),
|
||||||
|
)
|
||||||
|
generated: List[Path] = []
|
||||||
|
for name, create_chart in charts:
|
||||||
|
path = output_dir / f"{name}_{timestamp}.png"
|
||||||
|
create_chart(results, path)
|
||||||
|
generated.append(path)
|
||||||
|
return generated
|
||||||
|
|
||||||
# 1. 首次延迟对比图
|
def _save_line_chart(
|
||||||
if asr_results or tts_results:
|
|
||||||
path = output_dir / f"first_latency_{timestamp}.png"
|
|
||||||
self._generate_first_latency_chart(asr_results, tts_results, path)
|
|
||||||
generated_files.append(path)
|
|
||||||
|
|
||||||
# 2. RTF 对比图
|
|
||||||
if asr_results or tts_results:
|
|
||||||
path = output_dir / f"rtf_{timestamp}.png"
|
|
||||||
self._generate_rtf_chart(asr_results, tts_results, path)
|
|
||||||
generated_files.append(path)
|
|
||||||
|
|
||||||
# 3. 吞吐量对比图
|
|
||||||
if asr_results or tts_results:
|
|
||||||
path = output_dir / f"throughput_{timestamp}.png"
|
|
||||||
self._generate_throughput_chart(asr_results, tts_results, path)
|
|
||||||
generated_files.append(path)
|
|
||||||
|
|
||||||
# 4. 总时间对比图
|
|
||||||
if asr_results or tts_results:
|
|
||||||
path = output_dir / f"total_time_{timestamp}.png"
|
|
||||||
self._generate_total_time_chart(asr_results, tts_results, path)
|
|
||||||
generated_files.append(path)
|
|
||||||
|
|
||||||
return generated_files
|
|
||||||
|
|
||||||
def _generate_first_latency_chart(
|
|
||||||
self,
|
self,
|
||||||
asr_results: List[AggregatedMetrics],
|
results: List[AggregatedMetrics],
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
output_path: Path,
|
output_path: Path,
|
||||||
|
*,
|
||||||
|
title: str,
|
||||||
|
ylabel: str,
|
||||||
|
avg_value: str,
|
||||||
|
p95_value: str,
|
||||||
|
reference_line: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""生成首次延迟对比图"""
|
|
||||||
_fig, ax = plt.subplots(figsize=(10, 6))
|
_fig, ax = plt.subplots(figsize=(10, 6))
|
||||||
|
levels = [result.concurrency_level for result in results]
|
||||||
|
avg_values = [getattr(result, avg_value) for result in results]
|
||||||
|
p95_values = [getattr(result, p95_value) for result in results]
|
||||||
|
|
||||||
levels = []
|
ax.plot(levels, avg_values, "o-", color=self._COLOR, label="平均值", linewidth=2, markersize=8)
|
||||||
if asr_results:
|
ax.plot(levels, p95_values, "s--", color=self._COLOR, label="P95", linewidth=1.5, markersize=6, alpha=0.7)
|
||||||
levels = [r.concurrency_level for r in asr_results]
|
if reference_line:
|
||||||
avg_values = [r.first_latency_avg for r in asr_results]
|
ax.axhline(y=1.0, color="red", linestyle=":", linewidth=1.5, label="RTF = 1.0 (实时)")
|
||||||
p95_values = [r.first_latency_p95 for r in asr_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["asr"],
|
ax.set_xlabel("并发数", fontsize=12)
|
||||||
label='ASR 首次响应 (Avg)', linewidth=2, markersize=8)
|
ax.set_ylabel(ylabel, fontsize=12)
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["asr"],
|
ax.set_title(title, fontsize=14, fontweight="bold")
|
||||||
label='ASR 首次响应 (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
if tts_results:
|
|
||||||
levels = [r.concurrency_level for r in tts_results]
|
|
||||||
avg_values = [r.first_latency_avg for r in tts_results]
|
|
||||||
p95_values = [r.first_latency_p95 for r in tts_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["tts"],
|
|
||||||
label='TTS 首包延迟 (Avg)', linewidth=2, markersize=8)
|
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["tts"],
|
|
||||||
label='TTS 首包延迟 (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
ax.set_xlabel('并发数', fontsize=12)
|
|
||||||
ax.set_ylabel('延迟 (ms)', fontsize=12)
|
|
||||||
ax.set_title('首次响应延迟 vs 并发数', fontsize=14, fontweight='bold')
|
|
||||||
ax.legend(loc='best')
|
|
||||||
ax.grid(True, alpha=0.3)
|
|
||||||
if levels:
|
|
||||||
ax.set_xticks(levels)
|
ax.set_xticks(levels)
|
||||||
|
ax.legend(loc="best")
|
||||||
plt.tight_layout()
|
|
||||||
plt.savefig(output_path, dpi=150)
|
|
||||||
plt.close()
|
|
||||||
|
|
||||||
def _generate_rtf_chart(
|
|
||||||
self,
|
|
||||||
asr_results: List[AggregatedMetrics],
|
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
output_path: Path,
|
|
||||||
) -> None:
|
|
||||||
"""生成 RTF 对比图"""
|
|
||||||
_fig, ax = plt.subplots(figsize=(10, 6))
|
|
||||||
|
|
||||||
if asr_results:
|
|
||||||
levels = [r.concurrency_level for r in asr_results]
|
|
||||||
avg_values = [r.rtf_avg for r in asr_results]
|
|
||||||
p95_values = [r.rtf_p95 for r in asr_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["asr"],
|
|
||||||
label='ASR RTF (Avg)', linewidth=2, markersize=8)
|
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["asr"],
|
|
||||||
label='ASR RTF (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
if tts_results:
|
|
||||||
levels = [r.concurrency_level for r in tts_results]
|
|
||||||
avg_values = [r.rtf_avg for r in tts_results]
|
|
||||||
p95_values = [r.rtf_p95 for r in tts_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["tts"],
|
|
||||||
label='TTS RTF (Avg)', linewidth=2, markersize=8)
|
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["tts"],
|
|
||||||
label='TTS RTF (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
# 添加 RTF=1.0 参考线
|
|
||||||
all_levels = set()
|
|
||||||
if asr_results:
|
|
||||||
all_levels.update(r.concurrency_level for r in asr_results)
|
|
||||||
if tts_results:
|
|
||||||
all_levels.update(r.concurrency_level for r in tts_results)
|
|
||||||
if all_levels:
|
|
||||||
ax.axhline(y=1.0, color='red', linestyle=':', linewidth=1.5,
|
|
||||||
label='RTF = 1.0 (实时)')
|
|
||||||
|
|
||||||
ax.set_xlabel('并发数', fontsize=12)
|
|
||||||
ax.set_ylabel('RTF', fontsize=12)
|
|
||||||
ax.set_title('RTF vs 并发数', fontsize=14, fontweight='bold')
|
|
||||||
ax.legend(loc='best')
|
|
||||||
ax.grid(True, alpha=0.3)
|
ax.grid(True, alpha=0.3)
|
||||||
|
|
||||||
plt.tight_layout()
|
plt.tight_layout()
|
||||||
plt.savefig(output_path, dpi=150)
|
plt.savefig(output_path, dpi=150)
|
||||||
plt.close()
|
plt.close()
|
||||||
|
|
||||||
def _generate_throughput_chart(
|
def _generate_latency_chart(self, results: List[AggregatedMetrics], output_path: Path) -> None:
|
||||||
self,
|
self._save_line_chart(
|
||||||
asr_results: List[AggregatedMetrics],
|
results,
|
||||||
tts_results: List[AggregatedMetrics],
|
output_path,
|
||||||
output_path: Path,
|
title="ASR 首次响应延迟 vs 并发数",
|
||||||
) -> None:
|
ylabel="延迟 (ms)",
|
||||||
"""生成吞吐量柱状图"""
|
avg_value="first_latency_avg",
|
||||||
|
p95_value="first_latency_p95",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _generate_rtf_chart(self, results: List[AggregatedMetrics], output_path: Path) -> None:
|
||||||
|
self._save_line_chart(
|
||||||
|
results,
|
||||||
|
output_path,
|
||||||
|
title="ASR RTF vs 并发数",
|
||||||
|
ylabel="RTF",
|
||||||
|
avg_value="rtf_avg",
|
||||||
|
p95_value="rtf_p95",
|
||||||
|
reference_line=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _generate_total_time_chart(self, results: List[AggregatedMetrics], output_path: Path) -> None:
|
||||||
|
self._save_line_chart(
|
||||||
|
results,
|
||||||
|
output_path,
|
||||||
|
title="ASR 总处理时间 vs 并发数",
|
||||||
|
ylabel="时间 (ms)",
|
||||||
|
avg_value="total_time_avg",
|
||||||
|
p95_value="total_time_p95",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _generate_throughput_chart(self, results: List[AggregatedMetrics], output_path: Path) -> None:
|
||||||
_fig, ax = plt.subplots(figsize=(10, 6))
|
_fig, ax = plt.subplots(figsize=(10, 6))
|
||||||
|
levels = [result.concurrency_level for result in results]
|
||||||
|
throughput = [result.throughput for result in results]
|
||||||
|
positions = np.arange(len(levels))
|
||||||
|
|
||||||
all_levels = sorted(set(
|
ax.bar(positions, throughput, 0.6, label="ASR", color=self._COLOR, alpha=0.8)
|
||||||
[r.concurrency_level for r in asr_results] +
|
ax.set_xlabel("并发数", fontsize=12)
|
||||||
[r.concurrency_level for r in tts_results]
|
ax.set_ylabel("吞吐量 (req/s)", fontsize=12)
|
||||||
))
|
ax.set_title("ASR 吞吐量 vs 并发数", fontsize=14, fontweight="bold")
|
||||||
|
ax.set_xticks(positions)
|
||||||
x = np.arange(len(all_levels))
|
ax.set_xticklabels([str(level) for level in levels])
|
||||||
width = 0.35
|
ax.legend(loc="best")
|
||||||
|
ax.grid(True, alpha=0.3, axis="y")
|
||||||
if asr_results:
|
|
||||||
asr_throughput = []
|
|
||||||
for level in all_levels:
|
|
||||||
r = next((r for r in asr_results if r.concurrency_level == level), None)
|
|
||||||
asr_throughput.append(r.throughput if r else 0)
|
|
||||||
ax.bar(x - width/2, asr_throughput, width, label='ASR',
|
|
||||||
color=self.colors["asr"], alpha=0.8)
|
|
||||||
|
|
||||||
if tts_results:
|
|
||||||
tts_throughput = []
|
|
||||||
for level in all_levels:
|
|
||||||
r = next((r for r in tts_results if r.concurrency_level == level), None)
|
|
||||||
tts_throughput.append(r.throughput if r else 0)
|
|
||||||
ax.bar(x + width/2, tts_throughput, width, label='TTS',
|
|
||||||
color=self.colors["tts"], alpha=0.8)
|
|
||||||
|
|
||||||
ax.set_xlabel('并发数', fontsize=12)
|
|
||||||
ax.set_ylabel('吞吐量 (req/s)', fontsize=12)
|
|
||||||
ax.set_title('吞吐量 vs 并发数', fontsize=14, fontweight='bold')
|
|
||||||
ax.set_xticks(x)
|
|
||||||
ax.set_xticklabels([str(level) for level in all_levels])
|
|
||||||
ax.legend(loc='best')
|
|
||||||
ax.grid(True, alpha=0.3, axis='y')
|
|
||||||
|
|
||||||
plt.tight_layout()
|
|
||||||
plt.savefig(output_path, dpi=150)
|
|
||||||
plt.close()
|
|
||||||
|
|
||||||
def _generate_total_time_chart(
|
|
||||||
self,
|
|
||||||
asr_results: List[AggregatedMetrics],
|
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
output_path: Path,
|
|
||||||
) -> None:
|
|
||||||
"""生成总时间对比图"""
|
|
||||||
_fig, ax = plt.subplots(figsize=(10, 6))
|
|
||||||
|
|
||||||
if asr_results:
|
|
||||||
levels = [r.concurrency_level for r in asr_results]
|
|
||||||
avg_values = [r.total_time_avg for r in asr_results]
|
|
||||||
p95_values = [r.total_time_p95 for r in asr_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["asr"],
|
|
||||||
label='ASR 总时间 (Avg)', linewidth=2, markersize=8)
|
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["asr"],
|
|
||||||
label='ASR 总时间 (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
if tts_results:
|
|
||||||
levels = [r.concurrency_level for r in tts_results]
|
|
||||||
avg_values = [r.total_time_avg for r in tts_results]
|
|
||||||
p95_values = [r.total_time_p95 for r in tts_results]
|
|
||||||
|
|
||||||
ax.plot(levels, avg_values, 'o-', color=self.colors["tts"],
|
|
||||||
label='TTS 总时间 (Avg)', linewidth=2, markersize=8)
|
|
||||||
ax.plot(levels, p95_values, 's--', color=self.colors["tts"],
|
|
||||||
label='TTS 总时间 (P95)', linewidth=1.5, markersize=6, alpha=0.7)
|
|
||||||
|
|
||||||
ax.set_xlabel('并发数', fontsize=12)
|
|
||||||
ax.set_ylabel('时间 (ms)', fontsize=12)
|
|
||||||
ax.set_title('总处理时间 vs 并发数', fontsize=14, fontweight='bold')
|
|
||||||
ax.legend(loc='best')
|
|
||||||
ax.grid(True, alpha=0.3)
|
|
||||||
|
|
||||||
plt.tight_layout()
|
plt.tight_layout()
|
||||||
plt.savefig(output_path, dpi=150)
|
plt.savefig(output_path, dpi=150)
|
||||||
plt.close()
|
plt.close()
|
||||||
|
|
|
||||||
|
|
@ -16,7 +16,6 @@ class MarkdownReporter:
|
||||||
def generate(
|
def generate(
|
||||||
self,
|
self,
|
||||||
asr_results: List[AggregatedMetrics],
|
asr_results: List[AggregatedMetrics],
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
output_path: Path,
|
output_path: Path,
|
||||||
config_info: Optional[dict] = None,
|
config_info: Optional[dict] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -25,7 +24,6 @@ class MarkdownReporter:
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
asr_results: ASR 测试结果
|
asr_results: ASR 测试结果
|
||||||
tts_results: TTS 测试结果
|
|
||||||
output_path: 输出文件路径
|
output_path: 输出文件路径
|
||||||
config_info: 配置信息
|
config_info: 配置信息
|
||||||
"""
|
"""
|
||||||
|
|
@ -48,12 +46,8 @@ class MarkdownReporter:
|
||||||
if asr_results:
|
if asr_results:
|
||||||
lines.extend(self._generate_asr_section(asr_results))
|
lines.extend(self._generate_asr_section(asr_results))
|
||||||
|
|
||||||
# TTS 结果
|
|
||||||
if tts_results:
|
|
||||||
lines.extend(self._generate_tts_section(tts_results))
|
|
||||||
|
|
||||||
# 结论
|
# 结论
|
||||||
lines.extend(self._generate_conclusions(asr_results, tts_results))
|
lines.extend(self._generate_conclusions(asr_results))
|
||||||
|
|
||||||
# 写入文件
|
# 写入文件
|
||||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
@ -101,52 +95,9 @@ class MarkdownReporter:
|
||||||
lines.append("")
|
lines.append("")
|
||||||
return lines
|
return lines
|
||||||
|
|
||||||
def _generate_tts_section(self, results: List[AggregatedMetrics]) -> List[str]:
|
|
||||||
"""生成 TTS 结果部分"""
|
|
||||||
lines = []
|
|
||||||
lines.append("## TTS 性能测试结果")
|
|
||||||
lines.append("")
|
|
||||||
|
|
||||||
# 延迟指标表格
|
|
||||||
lines.append("### 延迟指标 (毫秒)")
|
|
||||||
lines.append("")
|
|
||||||
lines.append("| 并发数 | 首包延迟 (Avg) | 首包延迟 (P95) | 总时间 (Avg) | 总时间 (P95) | 总时间 (Max) |")
|
|
||||||
lines.append("|--------|---------------|---------------|-------------|-------------|-------------|")
|
|
||||||
|
|
||||||
for r in results:
|
|
||||||
lines.append(
|
|
||||||
f"| {r.concurrency_level} | "
|
|
||||||
f"{r.first_latency_avg:.1f} | "
|
|
||||||
f"{r.first_latency_p95:.1f} | "
|
|
||||||
f"{r.total_time_avg:.1f} | "
|
|
||||||
f"{r.total_time_p95:.1f} | "
|
|
||||||
f"{r.total_time_max:.1f} |"
|
|
||||||
)
|
|
||||||
|
|
||||||
lines.append("")
|
|
||||||
|
|
||||||
# RTF 和吞吐量表格
|
|
||||||
lines.append("### RTF 和吞吐量")
|
|
||||||
lines.append("")
|
|
||||||
lines.append("| 并发数 | RTF (Avg) | RTF (P95) | 吞吐量 (req/s) | 成功率 |")
|
|
||||||
lines.append("|--------|----------|----------|---------------|--------|")
|
|
||||||
|
|
||||||
for r in results:
|
|
||||||
lines.append(
|
|
||||||
f"| {r.concurrency_level} | "
|
|
||||||
f"{r.rtf_avg:.3f} | "
|
|
||||||
f"{r.rtf_p95:.3f} | "
|
|
||||||
f"{r.throughput:.3f} | "
|
|
||||||
f"{r.success_rate:.1f}% |"
|
|
||||||
)
|
|
||||||
|
|
||||||
lines.append("")
|
|
||||||
return lines
|
|
||||||
|
|
||||||
def _generate_conclusions(
|
def _generate_conclusions(
|
||||||
self,
|
self,
|
||||||
asr_results: List[AggregatedMetrics],
|
asr_results: List[AggregatedMetrics],
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""生成结论部分"""
|
"""生成结论部分"""
|
||||||
lines = []
|
lines = []
|
||||||
|
|
@ -164,16 +115,6 @@ class MarkdownReporter:
|
||||||
max_stable = max(stable_levels, key=lambda x: x.concurrency_level)
|
max_stable = max(stable_levels, key=lambda x: x.concurrency_level)
|
||||||
lines.append(f"- **ASR 稳定并发上限 (RTF < 1.0):** {max_stable.concurrency_level}")
|
lines.append(f"- **ASR 稳定并发上限 (RTF < 1.0):** {max_stable.concurrency_level}")
|
||||||
|
|
||||||
if tts_results:
|
|
||||||
max_level = max(tts_results, key=lambda x: x.concurrency_level)
|
|
||||||
lines.append(f"- **TTS 最大并发 ({max_level.concurrency_level}) RTF:** {max_level.rtf_avg:.3f}")
|
|
||||||
lines.append(f"- **TTS 最大并发吞吐量:** {max_level.throughput:.3f} req/s")
|
|
||||||
|
|
||||||
stable_levels = [r for r in tts_results if r.rtf_avg <= 1.0]
|
|
||||||
if stable_levels:
|
|
||||||
max_stable = max(stable_levels, key=lambda x: x.concurrency_level)
|
|
||||||
lines.append(f"- **TTS 稳定并发上限 (RTF < 1.0):** {max_stable.concurrency_level}")
|
|
||||||
|
|
||||||
lines.append("")
|
lines.append("")
|
||||||
lines.append("---")
|
lines.append("---")
|
||||||
lines.append("")
|
lines.append("")
|
||||||
|
|
|
||||||
|
|
@ -3,15 +3,9 @@
|
||||||
Qwen3-ASR 并发性能测试主入口
|
Qwen3-ASR 并发性能测试主入口
|
||||||
|
|
||||||
使用方法:
|
使用方法:
|
||||||
# 完整测试 (ASR + TTS)
|
# ASR 并发测试
|
||||||
python -m scripts.benchmark.run --audio-file /path/to/audio.wav
|
python -m scripts.benchmark.run --audio-file /path/to/audio.wav
|
||||||
|
|
||||||
# 仅测试 TTS
|
|
||||||
python -m scripts.benchmark.run --test-type tts
|
|
||||||
|
|
||||||
# 仅测试 ASR
|
|
||||||
python -m scripts.benchmark.run --audio-file /path/to/audio.wav --test-type asr
|
|
||||||
|
|
||||||
# 自定义并发级别
|
# 自定义并发级别
|
||||||
python -m scripts.benchmark.run --audio-file /path/to/audio.wav --concurrency 5 10 20
|
python -m scripts.benchmark.run --audio-file /path/to/audio.wav --concurrency 5 10 20
|
||||||
"""
|
"""
|
||||||
|
|
@ -26,13 +20,11 @@ from typing import List
|
||||||
|
|
||||||
from .config import TestConfig
|
from .config import TestConfig
|
||||||
from .clients.asr_client import ASRWebSocketClient
|
from .clients.asr_client import ASRWebSocketClient
|
||||||
from .clients.tts_client import TTSWebSocketClient
|
from .metrics.models import ASRMetrics, AggregatedMetrics
|
||||||
from .metrics.models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
|
||||||
from .metrics.statistics import calculate_statistics
|
from .metrics.statistics import calculate_statistics
|
||||||
from .reporters.markdown_reporter import MarkdownReporter
|
from .reporters.markdown_reporter import MarkdownReporter
|
||||||
from .reporters.chart_generator import ChartGenerator
|
from .reporters.chart_generator import ChartGenerator
|
||||||
from .utils.audio_utils import load_audio_file
|
from .utils.audio_utils import load_audio_file
|
||||||
from .utils.text_generator import generate_test_texts
|
|
||||||
|
|
||||||
# 配置日志
|
# 配置日志
|
||||||
logging.basicConfig(
|
logging.basicConfig(
|
||||||
|
|
@ -53,12 +45,8 @@ class ConcurrentBenchmark:
|
||||||
def _setup_output_dirs(self):
|
def _setup_output_dirs(self):
|
||||||
"""创建输出目录结构"""
|
"""创建输出目录结构"""
|
||||||
self.config.output_dir.mkdir(parents=True, exist_ok=True)
|
self.config.output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
# ASR 结果目录
|
|
||||||
self.asr_output_dir = self.config.output_dir / "asr"
|
self.asr_output_dir = self.config.output_dir / "asr"
|
||||||
self.asr_output_dir.mkdir(exist_ok=True)
|
self.asr_output_dir.mkdir(exist_ok=True)
|
||||||
# TTS 音频目录
|
|
||||||
self.tts_output_dir = self.config.output_dir / "tts"
|
|
||||||
self.tts_output_dir.mkdir(exist_ok=True)
|
|
||||||
|
|
||||||
async def run_asr_benchmark(self) -> List[AggregatedMetrics]:
|
async def run_asr_benchmark(self) -> List[AggregatedMetrics]:
|
||||||
"""
|
"""
|
||||||
|
|
@ -159,120 +147,15 @@ class ConcurrentBenchmark:
|
||||||
|
|
||||||
return metrics_list
|
return metrics_list
|
||||||
|
|
||||||
async def run_tts_benchmark(self) -> List[AggregatedMetrics]:
|
|
||||||
"""
|
|
||||||
运行 TTS 并发测试
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
各并发级别的聚合指标列表
|
|
||||||
"""
|
|
||||||
logger.info("开始 TTS 并发性能测试...")
|
|
||||||
|
|
||||||
# 生成测试文本
|
|
||||||
test_texts = generate_test_texts(
|
|
||||||
count=self.config.tts_text_count,
|
|
||||||
length_range=self.config.tts_text_length_range,
|
|
||||||
)
|
|
||||||
logger.info(f"已生成 {len(test_texts)} 段测试文本")
|
|
||||||
|
|
||||||
results = []
|
|
||||||
|
|
||||||
for level in self.config.concurrency_levels:
|
|
||||||
logger.info(f"\n测试并发级别: {level}")
|
|
||||||
|
|
||||||
# 选择文本 (每个并发请求使用不同文本)
|
|
||||||
selected_texts = test_texts[:level]
|
|
||||||
|
|
||||||
# 预热
|
|
||||||
logger.info(f" 预热中 ({min(self.config.warmup_requests, level)} 次请求)...")
|
|
||||||
await self._run_tts_concurrent(
|
|
||||||
selected_texts[:min(self.config.warmup_requests, level)],
|
|
||||||
min(self.config.warmup_requests, level),
|
|
||||||
level,
|
|
||||||
save_audio=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 正式测试
|
|
||||||
logger.info(f" 正式测试中...")
|
|
||||||
start_time = time.perf_counter()
|
|
||||||
metrics_list = await self._run_tts_concurrent(
|
|
||||||
selected_texts, level, level,
|
|
||||||
save_audio=True, # 正式测试时保存音频
|
|
||||||
)
|
|
||||||
total_time = time.perf_counter() - start_time
|
|
||||||
|
|
||||||
# 统计
|
|
||||||
aggregated = calculate_statistics(metrics_list, level, total_time)
|
|
||||||
results.append(aggregated)
|
|
||||||
|
|
||||||
# 打印结果
|
|
||||||
logger.info(f" 完成: 成功 {aggregated.successful_requests}/{aggregated.total_requests}")
|
|
||||||
logger.info(f" 首包延迟: {aggregated.first_latency_avg:.1f} ms (avg)")
|
|
||||||
logger.info(f" RTF: {aggregated.rtf_avg:.3f} (avg)")
|
|
||||||
|
|
||||||
return results
|
|
||||||
|
|
||||||
async def _run_tts_concurrent(
|
|
||||||
self,
|
|
||||||
texts: List[str],
|
|
||||||
num_requests: int,
|
|
||||||
concurrency_level: int,
|
|
||||||
save_audio: bool = False,
|
|
||||||
) -> List[TTSMetrics]:
|
|
||||||
"""运行并发 TTS 请求"""
|
|
||||||
tasks = []
|
|
||||||
|
|
||||||
for i in range(num_requests):
|
|
||||||
text = texts[i % len(texts)]
|
|
||||||
# 第一个请求始终开启调试模式
|
|
||||||
debug = (i == 0)
|
|
||||||
client = TTSWebSocketClient(
|
|
||||||
ws_url=self.config.tts_ws_url,
|
|
||||||
text=text,
|
|
||||||
voice=self.config.tts_voice,
|
|
||||||
audio_format=self.config.tts_format,
|
|
||||||
sample_rate=self.config.tts_sample_rate,
|
|
||||||
timeout=self.config.timeout_seconds,
|
|
||||||
chunk_interval=self.config.tts_chunk_interval,
|
|
||||||
debug=debug,
|
|
||||||
save_audio_dir=self.tts_output_dir if save_audio else None,
|
|
||||||
)
|
|
||||||
tasks.append(client.run_test())
|
|
||||||
|
|
||||||
# 添加进度提示
|
|
||||||
logger.info(f" 启动 {num_requests} 个并发请求...")
|
|
||||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
|
||||||
logger.info(f" 所有请求已完成")
|
|
||||||
|
|
||||||
# 处理结果
|
|
||||||
metrics_list = []
|
|
||||||
for result in results:
|
|
||||||
if isinstance(result, TTSMetrics):
|
|
||||||
result.concurrency_level = concurrency_level
|
|
||||||
metrics_list.append(result)
|
|
||||||
else:
|
|
||||||
# 异常情况
|
|
||||||
metrics = TTSMetrics(
|
|
||||||
request_id="error",
|
|
||||||
concurrency_level=concurrency_level,
|
|
||||||
start_time=0,
|
|
||||||
error_message=str(result),
|
|
||||||
)
|
|
||||||
metrics_list.append(metrics)
|
|
||||||
|
|
||||||
return metrics_list
|
|
||||||
|
|
||||||
def generate_report(
|
def generate_report(
|
||||||
self,
|
self,
|
||||||
asr_results: List[AggregatedMetrics],
|
asr_results: List[AggregatedMetrics],
|
||||||
tts_results: List[AggregatedMetrics],
|
|
||||||
) -> Path:
|
) -> Path:
|
||||||
"""
|
"""
|
||||||
生成测试报告
|
生成测试报告
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
asr_results: ASR 测试结果
|
asr_results: ASR 测试结果
|
||||||
tts_results: TTS 测试结果
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
报告文件路径
|
报告文件路径
|
||||||
|
|
@ -292,13 +175,13 @@ class ConcurrentBenchmark:
|
||||||
# 生成 Markdown 报告
|
# 生成 Markdown 报告
|
||||||
report_path = output_dir / f"{self.config.report_name}_{timestamp}.md"
|
report_path = output_dir / f"{self.config.report_name}_{timestamp}.md"
|
||||||
reporter = MarkdownReporter()
|
reporter = MarkdownReporter()
|
||||||
reporter.generate(asr_results, tts_results, report_path, config_info)
|
reporter.generate(asr_results, report_path, config_info)
|
||||||
logger.info(f"Markdown 报告已生成: {report_path}")
|
logger.info(f"Markdown 报告已生成: {report_path}")
|
||||||
|
|
||||||
# 生成图表
|
# 生成图表
|
||||||
chart_generator = ChartGenerator()
|
chart_generator = ChartGenerator()
|
||||||
chart_files = chart_generator.generate_all_charts(
|
chart_files = chart_generator.generate_all_charts(
|
||||||
asr_results, tts_results, output_dir, timestamp
|
asr_results, output_dir, timestamp
|
||||||
)
|
)
|
||||||
for chart_file in chart_files:
|
for chart_file in chart_files:
|
||||||
logger.info(f"图表已生成: {chart_file}")
|
logger.info(f"图表已生成: {chart_file}")
|
||||||
|
|
@ -313,12 +196,8 @@ def parse_args():
|
||||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||||
epilog="""
|
epilog="""
|
||||||
示例:
|
示例:
|
||||||
# 完整测试 (ASR + TTS)
|
|
||||||
python -m scripts.benchmark.run --audio-file test.wav
|
python -m scripts.benchmark.run --audio-file test.wav
|
||||||
|
|
||||||
# 仅测试 TTS
|
|
||||||
python -m scripts.benchmark.run --test-type tts
|
|
||||||
|
|
||||||
# 自定义并发级别
|
# 自定义并发级别
|
||||||
python -m scripts.benchmark.run --audio-file test.wav --concurrency 5 10 20 50
|
python -m scripts.benchmark.run --audio-file test.wav --concurrency 5 10 20 50
|
||||||
""",
|
""",
|
||||||
|
|
@ -338,7 +217,7 @@ def parse_args():
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--audio-file",
|
"--audio-file",
|
||||||
type=Path,
|
type=Path,
|
||||||
help="ASR 测试用音频文件路径 (测试 ASR 时必需)",
|
help="ASR 测试用音频文件路径(必需)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--concurrency",
|
"--concurrency",
|
||||||
|
|
@ -347,12 +226,6 @@ def parse_args():
|
||||||
default=[5, 10, 20, 50],
|
default=[5, 10, 20, 50],
|
||||||
help="并发级别列表 (默认: 5 10 20 50)",
|
help="并发级别列表 (默认: 5 10 20 50)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
|
||||||
"--test-type",
|
|
||||||
choices=["asr", "tts", "both"],
|
|
||||||
default="both",
|
|
||||||
help="测试类型 (默认: both)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--output",
|
"--output",
|
||||||
type=Path,
|
type=Path,
|
||||||
|
|
@ -365,12 +238,6 @@ def parse_args():
|
||||||
default=120.0,
|
default=120.0,
|
||||||
help="请求超时时间 (秒, 默认: 120)",
|
help="请求超时时间 (秒, 默认: 120)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
|
||||||
"--voice",
|
|
||||||
default="中文女",
|
|
||||||
help="TTS 音色 (默认: 中文女)",
|
|
||||||
)
|
|
||||||
|
|
||||||
return parser.parse_args()
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -386,12 +253,11 @@ async def main():
|
||||||
asr_audio_file=args.audio_file,
|
asr_audio_file=args.audio_file,
|
||||||
output_dir=args.output,
|
output_dir=args.output,
|
||||||
timeout_seconds=args.timeout,
|
timeout_seconds=args.timeout,
|
||||||
tts_voice=args.voice,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 验证配置
|
# 验证配置
|
||||||
try:
|
try:
|
||||||
config.validate(args.test_type)
|
config.validate()
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
logger.error(f"配置错误: {e}")
|
logger.error(f"配置错误: {e}")
|
||||||
return
|
return
|
||||||
|
|
@ -399,18 +265,10 @@ async def main():
|
||||||
# 运行测试
|
# 运行测试
|
||||||
benchmark = ConcurrentBenchmark(config)
|
benchmark = ConcurrentBenchmark(config)
|
||||||
|
|
||||||
asr_results = []
|
|
||||||
tts_results = []
|
|
||||||
|
|
||||||
if args.test_type in ("asr", "both"):
|
|
||||||
asr_results = await benchmark.run_asr_benchmark()
|
asr_results = await benchmark.run_asr_benchmark()
|
||||||
|
|
||||||
if args.test_type in ("tts", "both"):
|
|
||||||
tts_results = await benchmark.run_tts_benchmark()
|
|
||||||
|
|
||||||
# 生成报告
|
# 生成报告
|
||||||
if asr_results or tts_results:
|
report_path = benchmark.generate_report(asr_results)
|
||||||
report_path = benchmark.generate_report(asr_results, tts_results)
|
|
||||||
logger.info(f"\n测试完成! 报告已保存到: {report_path}")
|
logger.info(f"\n测试完成! 报告已保存到: {report_path}")
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,4 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
from .audio_utils import load_audio_file, get_audio_duration
|
from .audio_utils import load_audio_file, get_audio_duration
|
||||||
from .text_generator import generate_test_texts
|
|
||||||
|
|
||||||
__all__ = ["load_audio_file", "get_audio_duration", "generate_test_texts"]
|
__all__ = ["load_audio_file", "get_audio_duration"]
|
||||||
|
|
|
||||||
|
|
@ -1,170 +0,0 @@
|
||||||
# -*- coding: utf-8 -*-
|
|
||||||
"""
|
|
||||||
中文随机句子生成器
|
|
||||||
|
|
||||||
用于生成 TTS 测试文本,不依赖外部 AI API。
|
|
||||||
"""
|
|
||||||
|
|
||||||
import random
|
|
||||||
from typing import List, Tuple
|
|
||||||
|
|
||||||
# 主语词库
|
|
||||||
SUBJECTS = [
|
|
||||||
"我", "你", "他", "她", "我们", "大家", "小明", "小红", "老师", "学生",
|
|
||||||
"医生", "工程师", "科学家", "艺术家", "音乐家", "作家", "记者", "警察",
|
|
||||||
"这位先生", "那位女士", "我的朋友", "他的同事", "她的家人", "公司",
|
|
||||||
"团队", "项目组", "研发部门", "市场部", "客户", "用户",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 时间词库
|
|
||||||
TIME_PHRASES = [
|
|
||||||
"今天", "明天", "昨天", "上周", "下周", "这个月", "上个月", "今年",
|
|
||||||
"最近", "刚才", "马上", "立刻", "很快", "不久前", "过去",
|
|
||||||
"早上", "中午", "下午", "晚上", "凌晨", "周末", "假期期间",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 地点词库
|
|
||||||
LOCATIONS = [
|
|
||||||
"在公司", "在家里", "在学校", "在图书馆", "在咖啡厅", "在会议室",
|
|
||||||
"在公园", "在商场", "在医院", "在机场", "在火车站", "在地铁站",
|
|
||||||
"在办公室", "在实验室", "在教室", "在操场", "在餐厅", "在酒店",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 动词短语词库
|
|
||||||
VERB_PHRASES = [
|
|
||||||
"正在开发一个新的功能", "完成了一项重要的任务", "参加了一个技术会议",
|
|
||||||
"学习了新的编程语言", "解决了一个复杂的问题", "提交了项目报告",
|
|
||||||
"设计了一套新的方案", "测试了最新的版本", "优化了系统性能",
|
|
||||||
"讨论了未来的发展计划", "制定了下一步的工作安排", "回顾了过去的工作成果",
|
|
||||||
"分析了市场数据", "研究了用户需求", "改进了产品体验",
|
|
||||||
"组织了团队活动", "培训了新员工", "更新了技术文档",
|
|
||||||
"修复了几个重要的问题", "部署了新的服务", "监控了系统运行状态",
|
|
||||||
"收集了用户反馈", "整理了项目资料", "准备了演示材料",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 形容词词库
|
|
||||||
ADJECTIVES = [
|
|
||||||
"高效的", "专业的", "创新的", "稳定的", "可靠的", "智能的",
|
|
||||||
"先进的", "实用的", "便捷的", "优秀的", "杰出的", "卓越的",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 名词词库
|
|
||||||
NOUNS = [
|
|
||||||
"系统", "平台", "应用", "服务", "方案", "产品", "技术", "工具",
|
|
||||||
"项目", "团队", "计划", "目标", "成果", "进展", "效率", "质量",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 连接词
|
|
||||||
CONNECTORS = [
|
|
||||||
"并且", "同时", "而且", "另外", "此外", "因此", "所以", "然后",
|
|
||||||
]
|
|
||||||
|
|
||||||
# 结尾语
|
|
||||||
ENDINGS = [
|
|
||||||
"这是一个很好的开始。",
|
|
||||||
"我们对此感到非常满意。",
|
|
||||||
"期待能有更好的结果。",
|
|
||||||
"这将带来积极的影响。",
|
|
||||||
"相信未来会更加美好。",
|
|
||||||
"让我们继续努力。",
|
|
||||||
"这是值得庆祝的成就。",
|
|
||||||
"我们会继续保持这种势头。",
|
|
||||||
"这体现了团队的实力。",
|
|
||||||
"我们为此感到自豪。",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def generate_simple_sentence() -> str:
|
|
||||||
"""生成简单句"""
|
|
||||||
subject = random.choice(SUBJECTS)
|
|
||||||
time_phrase = random.choice(TIME_PHRASES) if random.random() > 0.3 else ""
|
|
||||||
location = random.choice(LOCATIONS) if random.random() > 0.5 else ""
|
|
||||||
verb_phrase = random.choice(VERB_PHRASES)
|
|
||||||
|
|
||||||
parts = [time_phrase, subject, location, verb_phrase]
|
|
||||||
parts = [p for p in parts if p] # 过滤空字符串
|
|
||||||
return "".join(parts) + "。"
|
|
||||||
|
|
||||||
|
|
||||||
def generate_compound_sentence() -> str:
|
|
||||||
"""生成复合句"""
|
|
||||||
sentence1 = generate_simple_sentence().rstrip("。")
|
|
||||||
connector = random.choice(CONNECTORS)
|
|
||||||
sentence2 = generate_simple_sentence().rstrip("。")
|
|
||||||
|
|
||||||
return f"{sentence1},{connector}{sentence2}。"
|
|
||||||
|
|
||||||
|
|
||||||
def generate_descriptive_sentence() -> str:
|
|
||||||
"""生成描述性句子"""
|
|
||||||
subject = random.choice(SUBJECTS)
|
|
||||||
adj = random.choice(ADJECTIVES)
|
|
||||||
noun = random.choice(NOUNS)
|
|
||||||
verb_phrase = random.choice(VERB_PHRASES)
|
|
||||||
|
|
||||||
return f"{subject}开发了一个{adj}{noun},{verb_phrase}。"
|
|
||||||
|
|
||||||
|
|
||||||
def generate_single_text(length_range: Tuple[int, int] = (50, 100)) -> str:
|
|
||||||
"""
|
|
||||||
生成单个测试文本
|
|
||||||
|
|
||||||
Args:
|
|
||||||
length_range: 文本长度范围 (min, max)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
生成的文本
|
|
||||||
"""
|
|
||||||
min_len, max_len = length_range
|
|
||||||
target_len = random.randint(min_len, max_len)
|
|
||||||
|
|
||||||
text = ""
|
|
||||||
sentence_generators = [
|
|
||||||
generate_simple_sentence,
|
|
||||||
generate_compound_sentence,
|
|
||||||
generate_descriptive_sentence,
|
|
||||||
]
|
|
||||||
|
|
||||||
while len(text) < target_len:
|
|
||||||
generator = random.choice(sentence_generators)
|
|
||||||
sentence = generator()
|
|
||||||
text += sentence
|
|
||||||
|
|
||||||
# 如果超出太多,截断到最近的句号
|
|
||||||
if len(text) > max_len + 20:
|
|
||||||
# 找到目标长度附近的句号
|
|
||||||
end_pos = text.rfind("。", 0, max_len + 10)
|
|
||||||
if end_pos > min_len:
|
|
||||||
text = text[: end_pos + 1]
|
|
||||||
|
|
||||||
return text
|
|
||||||
|
|
||||||
|
|
||||||
def generate_test_texts(
|
|
||||||
count: int = 50,
|
|
||||||
length_range: Tuple[int, int] = (50, 100),
|
|
||||||
) -> List[str]:
|
|
||||||
"""
|
|
||||||
生成测试文本列表
|
|
||||||
|
|
||||||
Args:
|
|
||||||
count: 生成数量
|
|
||||||
length_range: 文本长度范围
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
文本列表
|
|
||||||
"""
|
|
||||||
texts = []
|
|
||||||
for _ in range(count):
|
|
||||||
text = generate_single_text(length_range)
|
|
||||||
texts.append(text)
|
|
||||||
|
|
||||||
return texts
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
# 测试文本生成
|
|
||||||
texts = generate_test_texts(5, (50, 100))
|
|
||||||
for i, text in enumerate(texts, 1):
|
|
||||||
print(f"[{i}] ({len(text)}字): {text}")
|
|
||||||
print()
|
|
||||||
Loading…
Reference in New Issue