移除TTS相关代码
parent
5b4ab416e2
commit
86c9f1d415
|
|
@ -1,6 +1,6 @@
|
|||
# Qwen3-ASR 并发性能测试脚本
|
||||
|
||||
测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。
|
||||
测试 ASR WebSocket 服务在不同并发级别下的性能表现。
|
||||
|
||||
## 依赖
|
||||
|
||||
|
|
@ -21,15 +21,8 @@ python start.py
|
|||
### 2. 运行测试
|
||||
|
||||
```bash
|
||||
# 完整测试 (ASR + TTS)
|
||||
.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 分段)
|
||||
.venv/bin/python -m scripts.benchmark.qwen_rust_sensitivity \
|
||||
--audio-file /path/to/audio.wav
|
||||
|
|
@ -41,12 +34,10 @@ python start.py
|
|||
|------|--------|------|
|
||||
| `--host` | localhost | 服务器主机名 |
|
||||
| `--port` | 8000 | 服务器端口 |
|
||||
| `--audio-file` | - | ASR 测试音频文件路径 (测试 ASR 时必需) |
|
||||
| `--test-type` | both | 测试类型: `asr` / `tts` / `both` |
|
||||
| `--audio-file` | - | ASR 测试音频文件路径(必需) |
|
||||
| `--concurrency` | 5 10 20 50 | 并发级别列表 |
|
||||
| `--output` | ./benchmark_results | 报告输出目录 |
|
||||
| `--timeout` | 120 | 请求超时时间 (秒) |
|
||||
| `--voice` | 中文女 | TTS 测试音色 |
|
||||
|
||||
## Qwen Rust CPU 固定配置跑测
|
||||
|
||||
|
|
@ -89,14 +80,6 @@ QWEN_RUST_ALIGN_CONCURRENCY=4 \
|
|||
--markdown-out temp/qwen_rust_runtime_config_report.md
|
||||
```
|
||||
|
||||
### TTS 流式模拟配置 (config.py)
|
||||
|
||||
TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、顿号等)分割文本逐步发送:
|
||||
|
||||
| 配置项 | 默认值 | 说明 |
|
||||
|--------|--------|------|
|
||||
| `tts_chunk_interval` | 0.05 | 发送间隔秒数 (模拟 LLM 生成速度) |
|
||||
|
||||
## 使用示例
|
||||
|
||||
```bash
|
||||
|
|
@ -105,16 +88,11 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
|||
--audio-file test.wav \
|
||||
--concurrency 5 10 20 50 100
|
||||
|
||||
# 连接远程服务器
|
||||
# 连接远程 ASR 服务
|
||||
.venv/bin/python -m scripts.benchmark.run \
|
||||
--host 192.168.1.100 \
|
||||
--port 8000 \
|
||||
--test-type tts
|
||||
|
||||
# 使用不同音色测试 TTS
|
||||
.venv/bin/python -m scripts.benchmark.run \
|
||||
--test-type tts \
|
||||
--voice 中文男
|
||||
--audio-file test.wav
|
||||
```
|
||||
|
||||
## 测试指标
|
||||
|
|
@ -124,11 +102,6 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
|||
- **总处理时间**: 从开始到识别完成的总时间
|
||||
- **RTF**: 处理时间 / 音频时长 (小于 1.0 表示快于实时)
|
||||
|
||||
### TTS 指标
|
||||
- **首包延迟**: 从发送文本到收到第一个音频块的时间
|
||||
- **总合成时间**: 从开始到合成完成的总时间
|
||||
- **RTF**: 合成时间 / 生成音频时长
|
||||
|
||||
### 统计维度
|
||||
每个指标计算: 平均值 (Avg)、P50、P95、P99、最大值 (Max)
|
||||
|
||||
|
|
@ -139,8 +112,8 @@ TTS 测试模拟 LLM 流式输出场景,按标点符号(逗号、句号、
|
|||
```
|
||||
benchmark_results/
|
||||
├── benchmark_report_20241202_143000.md # Markdown 报告
|
||||
├── first_latency_20241202_143000.png # 首次响应延迟图
|
||||
├── rtf_20241202_143000.png # RTF 对比图
|
||||
├── first_latency_20241202_143000.png # ASR 首次响应延迟图
|
||||
├── rtf_20241202_143000.png # ASR RTF 图
|
||||
├── throughput_20241202_143000.png # 吞吐量图
|
||||
└── total_time_20241202_143000.png # 总时间图
|
||||
```
|
||||
|
|
@ -174,7 +147,6 @@ scripts/benchmark/
|
|||
├── clients/
|
||||
│ ├── base_client.py # WebSocket 客户端基类
|
||||
│ ├── asr_client.py # ASR 测试客户端
|
||||
│ └── tts_client.py # TTS 测试客户端
|
||||
├── metrics/
|
||||
│ ├── models.py # 指标数据类
|
||||
│ └── statistics.py # 统计计算
|
||||
|
|
@ -182,17 +154,14 @@ scripts/benchmark/
|
|||
│ ├── markdown_reporter.py # Markdown 报告生成
|
||||
│ └── chart_generator.py # 图表生成
|
||||
└── utils/
|
||||
├── audio_utils.py # 音频文件处理
|
||||
└── text_generator.py # 测试文本生成
|
||||
└── audio_utils.py # 音频文件处理
|
||||
```
|
||||
|
||||
## 注意事项
|
||||
|
||||
1. **ASR 测试需要音频文件**: 建议使用 1 分钟左右的音频,格式支持 wav/mp3 等常见格式
|
||||
2. **TTS 测试自动生成文本**: 使用内置的中文随机句子生成器,无需额外准备
|
||||
3. **TTS 模拟流式输入**: 测试会按标点符号(逗号、句号、顿号等)分割文本逐步发送,模拟 LLM 流式输出场景
|
||||
4. **并发测试会占用资源**: 高并发测试时请确保服务器有足够资源
|
||||
5. **RTF 解读**:
|
||||
2. **并发测试会占用资源**: 高并发测试时请确保服务器有足够资源
|
||||
3. **RTF 解读**:
|
||||
- RTF < 1.0: 处理速度快于实时,性能良好
|
||||
- RTF ≈ 1.0: 刚好实时处理
|
||||
- RTF > 1.0: 处理速度慢于实时,可能出现延迟累积
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
"""
|
||||
Qwen3-ASR 并发性能测试脚本
|
||||
|
||||
用于测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。
|
||||
用于测试 ASR WebSocket 服务在不同并发级别下的性能表现。
|
||||
"""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
# -*- coding: utf-8 -*-
|
||||
from .base_client import BaseWebSocketClient
|
||||
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"
|
||||
port: int = 8000
|
||||
timeout_seconds: float = 300.0 # 默认 5 分钟,并发 TTS 可能需要更长时间
|
||||
timeout_seconds: float = 300.0
|
||||
warmup_requests: int = 3
|
||||
|
||||
# 并发配置
|
||||
|
|
@ -27,14 +27,6 @@ class TestConfig:
|
|||
asr_chunk_size: int = 9600 # 600ms @ 16kHz
|
||||
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"))
|
||||
report_name: str = "benchmark_report"
|
||||
|
|
@ -49,26 +41,17 @@ class TestConfig:
|
|||
"""ASR WebSocket URL"""
|
||||
return f"{self.ws_base_url}/ws/v1/asr"
|
||||
|
||||
@property
|
||||
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:
|
||||
def validate(self) -> None:
|
||||
"""
|
||||
验证配置
|
||||
|
||||
Args:
|
||||
test_type: 测试类型 (asr/tts/both)
|
||||
|
||||
Raises:
|
||||
ValueError: 配置无效
|
||||
"""
|
||||
if test_type in ("asr", "both"):
|
||||
if self.asr_audio_file is None:
|
||||
raise ValueError("ASR 测试需要提供音频文件路径 (--audio-file)")
|
||||
if not self.asr_audio_file.exists():
|
||||
raise ValueError(f"音频文件不存在: {self.asr_audio_file}")
|
||||
if self.asr_audio_file is None:
|
||||
raise ValueError("ASR 测试需要提供音频文件路径 (--audio-file)")
|
||||
if not self.asr_audio_file.exists():
|
||||
raise ValueError(f"音频文件不存在: {self.asr_audio_file}")
|
||||
|
||||
if not self.concurrency_levels:
|
||||
raise ValueError("至少需要一个并发级别")
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
# -*- coding: utf-8 -*-
|
||||
from .models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
||||
from .models import ASRMetrics, AggregatedMetrics
|
||||
from .statistics import calculate_statistics, calculate_percentile
|
||||
|
||||
__all__ = [
|
||||
"ASRMetrics",
|
||||
"TTSMetrics",
|
||||
"AggregatedMetrics",
|
||||
"calculate_statistics",
|
||||
"calculate_percentile",
|
||||
|
|
|
|||
|
|
@ -49,64 +49,10 @@ class ASRMetrics:
|
|||
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
|
||||
class AggregatedMetrics:
|
||||
"""聚合后的指标 (针对一个并发级别)"""
|
||||
|
||||
test_type: str # "asr" or "tts"
|
||||
concurrency_level: int
|
||||
total_requests: int
|
||||
successful_requests: int
|
||||
|
|
|
|||
|
|
@ -3,10 +3,10 @@
|
|||
统计计算模块
|
||||
"""
|
||||
|
||||
from typing import List, Union
|
||||
from typing import List
|
||||
import numpy as np
|
||||
|
||||
from .models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
||||
from .models import ASRMetrics, AggregatedMetrics
|
||||
|
||||
|
||||
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))
|
||||
|
||||
|
||||
def calculate_asr_statistics(
|
||||
def calculate_statistics(
|
||||
metrics_list: List[ASRMetrics],
|
||||
concurrency_level: int,
|
||||
total_test_time: float,
|
||||
|
|
@ -41,6 +41,9 @@ def calculate_asr_statistics(
|
|||
Returns:
|
||||
聚合后的指标
|
||||
"""
|
||||
if not metrics_list:
|
||||
raise ValueError("指标列表不能为空")
|
||||
|
||||
successful = [m for m in metrics_list if 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]
|
||||
|
||||
return AggregatedMetrics(
|
||||
test_type="asr",
|
||||
concurrency_level=concurrency_level,
|
||||
total_requests=len(metrics_list),
|
||||
successful_requests=len(successful),
|
||||
|
|
@ -79,90 +81,3 @@ def calculate_asr_statistics(
|
|||
rtf_p99=calculate_percentile(rtfs, 99),
|
||||
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 -*-
|
||||
"""
|
||||
Matplotlib 图表生成器
|
||||
"""
|
||||
"""Generate ASR benchmark charts."""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
from ..metrics.models import AggregatedMetrics
|
||||
|
||||
# 设置中文字体支持
|
||||
matplotlib.rcParams['font.sans-serif'] = ['Arial Unicode MS', 'SimHei', 'DejaVu Sans']
|
||||
matplotlib.rcParams['axes.unicode_minus'] = False
|
||||
# Keep generated labels readable on common Windows and Linux installations.
|
||||
matplotlib.rcParams["font.sans-serif"] = ["Arial Unicode MS", "SimHei", "DejaVu Sans"]
|
||||
matplotlib.rcParams["axes.unicode_minus"] = False
|
||||
|
||||
|
||||
class ChartGenerator:
|
||||
"""图表生成器"""
|
||||
"""Create the latency, RTF, throughput, and duration ASR charts."""
|
||||
|
||||
def __init__(self):
|
||||
self.colors = {
|
||||
"asr": "#4CAF50", # 绿色
|
||||
"tts": "#2196F3", # 蓝色
|
||||
}
|
||||
_COLOR = "#4CAF50"
|
||||
|
||||
def generate_all_charts(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
results: List[AggregatedMetrics],
|
||||
output_dir: Path,
|
||||
timestamp: str,
|
||||
) -> List[Path]:
|
||||
"""
|
||||
生成所有图表
|
||||
|
||||
Args:
|
||||
asr_results: ASR 测试结果
|
||||
tts_results: TTS 测试结果
|
||||
output_dir: 输出目录
|
||||
timestamp: 时间戳
|
||||
|
||||
Returns:
|
||||
生成的图表文件路径列表
|
||||
"""
|
||||
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. 首次延迟对比图
|
||||
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(
|
||||
def _save_line_chart(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
results: List[AggregatedMetrics],
|
||||
output_path: Path,
|
||||
*,
|
||||
title: str,
|
||||
ylabel: str,
|
||||
avg_value: str,
|
||||
p95_value: str,
|
||||
reference_line: bool = False,
|
||||
) -> None:
|
||||
"""生成首次延迟对比图"""
|
||||
_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 = []
|
||||
if asr_results:
|
||||
levels = [r.concurrency_level for r in asr_results]
|
||||
avg_values = [r.first_latency_avg for r in asr_results]
|
||||
p95_values = [r.first_latency_p95 for r in asr_results]
|
||||
ax.plot(levels, avg_values, "o-", color=self._COLOR, label="平均值", linewidth=2, markersize=8)
|
||||
ax.plot(levels, p95_values, "s--", color=self._COLOR, label="P95", linewidth=1.5, markersize=6, alpha=0.7)
|
||||
if reference_line:
|
||||
ax.axhline(y=1.0, color="red", linestyle=":", linewidth=1.5, label="RTF = 1.0 (实时)")
|
||||
|
||||
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.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.set_xlabel("并发数", fontsize=12)
|
||||
ax.set_ylabel(ylabel, fontsize=12)
|
||||
ax.set_title(title, fontsize=14, fontweight="bold")
|
||||
ax.set_xticks(levels)
|
||||
ax.legend(loc="best")
|
||||
ax.grid(True, alpha=0.3)
|
||||
if levels:
|
||||
ax.set_xticks(levels)
|
||||
|
||||
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 对比图"""
|
||||
def _generate_latency_chart(self, results: List[AggregatedMetrics], output_path: Path) -> None:
|
||||
self._save_line_chart(
|
||||
results,
|
||||
output_path,
|
||||
title="ASR 首次响应延迟 vs 并发数",
|
||||
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))
|
||||
levels = [result.concurrency_level for result in results]
|
||||
throughput = [result.throughput for result in results]
|
||||
positions = np.arange(len(levels))
|
||||
|
||||
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)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig(output_path, dpi=150)
|
||||
plt.close()
|
||||
|
||||
def _generate_throughput_chart(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
output_path: Path,
|
||||
) -> None:
|
||||
"""生成吞吐量柱状图"""
|
||||
_fig, ax = plt.subplots(figsize=(10, 6))
|
||||
|
||||
all_levels = sorted(set(
|
||||
[r.concurrency_level for r in asr_results] +
|
||||
[r.concurrency_level for r in tts_results]
|
||||
))
|
||||
|
||||
x = np.arange(len(all_levels))
|
||||
width = 0.35
|
||||
|
||||
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)
|
||||
|
||||
ax.bar(positions, throughput, 0.6, label="ASR", color=self._COLOR, alpha=0.8)
|
||||
ax.set_xlabel("并发数", fontsize=12)
|
||||
ax.set_ylabel("吞吐量 (req/s)", fontsize=12)
|
||||
ax.set_title("ASR 吞吐量 vs 并发数", fontsize=14, fontweight="bold")
|
||||
ax.set_xticks(positions)
|
||||
ax.set_xticklabels([str(level) for level in levels])
|
||||
ax.legend(loc="best")
|
||||
ax.grid(True, alpha=0.3, axis="y")
|
||||
plt.tight_layout()
|
||||
plt.savefig(output_path, dpi=150)
|
||||
plt.close()
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ class MarkdownReporter:
|
|||
def generate(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
output_path: Path,
|
||||
config_info: Optional[dict] = None,
|
||||
) -> None:
|
||||
|
|
@ -25,7 +24,6 @@ class MarkdownReporter:
|
|||
|
||||
Args:
|
||||
asr_results: ASR 测试结果
|
||||
tts_results: TTS 测试结果
|
||||
output_path: 输出文件路径
|
||||
config_info: 配置信息
|
||||
"""
|
||||
|
|
@ -48,12 +46,8 @@ class MarkdownReporter:
|
|||
if 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)
|
||||
|
|
@ -101,52 +95,9 @@ class MarkdownReporter:
|
|||
lines.append("")
|
||||
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(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
) -> List[str]:
|
||||
"""生成结论部分"""
|
||||
lines = []
|
||||
|
|
@ -164,16 +115,6 @@ class MarkdownReporter:
|
|||
max_stable = max(stable_levels, key=lambda x: x.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("")
|
||||
|
|
|
|||
|
|
@ -3,15 +3,9 @@
|
|||
Qwen3-ASR 并发性能测试主入口
|
||||
|
||||
使用方法:
|
||||
# 完整测试 (ASR + TTS)
|
||||
# ASR 并发测试
|
||||
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
|
||||
"""
|
||||
|
|
@ -26,13 +20,11 @@ from typing import List
|
|||
|
||||
from .config import TestConfig
|
||||
from .clients.asr_client import ASRWebSocketClient
|
||||
from .clients.tts_client import TTSWebSocketClient
|
||||
from .metrics.models import ASRMetrics, TTSMetrics, AggregatedMetrics
|
||||
from .metrics.models import ASRMetrics, AggregatedMetrics
|
||||
from .metrics.statistics import calculate_statistics
|
||||
from .reporters.markdown_reporter import MarkdownReporter
|
||||
from .reporters.chart_generator import ChartGenerator
|
||||
from .utils.audio_utils import load_audio_file
|
||||
from .utils.text_generator import generate_test_texts
|
||||
|
||||
# 配置日志
|
||||
logging.basicConfig(
|
||||
|
|
@ -53,12 +45,8 @@ class ConcurrentBenchmark:
|
|||
def _setup_output_dirs(self):
|
||||
"""创建输出目录结构"""
|
||||
self.config.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
# ASR 结果目录
|
||||
self.asr_output_dir = self.config.output_dir / "asr"
|
||||
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]:
|
||||
"""
|
||||
|
|
@ -159,120 +147,15 @@ class ConcurrentBenchmark:
|
|||
|
||||
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(
|
||||
self,
|
||||
asr_results: List[AggregatedMetrics],
|
||||
tts_results: List[AggregatedMetrics],
|
||||
) -> Path:
|
||||
"""
|
||||
生成测试报告
|
||||
|
||||
Args:
|
||||
asr_results: ASR 测试结果
|
||||
tts_results: TTS 测试结果
|
||||
|
||||
Returns:
|
||||
报告文件路径
|
||||
|
|
@ -292,13 +175,13 @@ class ConcurrentBenchmark:
|
|||
# 生成 Markdown 报告
|
||||
report_path = output_dir / f"{self.config.report_name}_{timestamp}.md"
|
||||
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}")
|
||||
|
||||
# 生成图表
|
||||
chart_generator = ChartGenerator()
|
||||
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:
|
||||
logger.info(f"图表已生成: {chart_file}")
|
||||
|
|
@ -313,12 +196,8 @@ def parse_args():
|
|||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
epilog="""
|
||||
示例:
|
||||
# 完整测试 (ASR + TTS)
|
||||
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
|
||||
""",
|
||||
|
|
@ -338,7 +217,7 @@ def parse_args():
|
|||
parser.add_argument(
|
||||
"--audio-file",
|
||||
type=Path,
|
||||
help="ASR 测试用音频文件路径 (测试 ASR 时必需)",
|
||||
help="ASR 测试用音频文件路径(必需)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--concurrency",
|
||||
|
|
@ -347,12 +226,6 @@ def parse_args():
|
|||
default=[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(
|
||||
"--output",
|
||||
type=Path,
|
||||
|
|
@ -365,12 +238,6 @@ def parse_args():
|
|||
default=120.0,
|
||||
help="请求超时时间 (秒, 默认: 120)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--voice",
|
||||
default="中文女",
|
||||
help="TTS 音色 (默认: 中文女)",
|
||||
)
|
||||
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
|
|
@ -386,12 +253,11 @@ async def main():
|
|||
asr_audio_file=args.audio_file,
|
||||
output_dir=args.output,
|
||||
timeout_seconds=args.timeout,
|
||||
tts_voice=args.voice,
|
||||
)
|
||||
|
||||
# 验证配置
|
||||
try:
|
||||
config.validate(args.test_type)
|
||||
config.validate()
|
||||
except ValueError as e:
|
||||
logger.error(f"配置错误: {e}")
|
||||
return
|
||||
|
|
@ -399,19 +265,11 @@ async def main():
|
|||
# 运行测试
|
||||
benchmark = ConcurrentBenchmark(config)
|
||||
|
||||
asr_results = []
|
||||
tts_results = []
|
||||
|
||||
if args.test_type in ("asr", "both"):
|
||||
asr_results = await benchmark.run_asr_benchmark()
|
||||
|
||||
if args.test_type in ("tts", "both"):
|
||||
tts_results = await benchmark.run_tts_benchmark()
|
||||
asr_results = await benchmark.run_asr_benchmark()
|
||||
|
||||
# 生成报告
|
||||
if asr_results or tts_results:
|
||||
report_path = benchmark.generate_report(asr_results, tts_results)
|
||||
logger.info(f"\n测试完成! 报告已保存到: {report_path}")
|
||||
report_path = benchmark.generate_report(asr_results)
|
||||
logger.info(f"\n测试完成! 报告已保存到: {report_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
# -*- coding: utf-8 -*-
|
||||
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