test/scripts/benchmark/run.py

277 lines
8.1 KiB
Python

# -*- coding: utf-8 -*-
"""
Qwen3-ASR 并发性能测试主入口
使用方法:
# ASR 并发测试
python -m scripts.benchmark.run --audio-file /path/to/audio.wav
# 自定义并发级别
python -m scripts.benchmark.run --audio-file /path/to/audio.wav --concurrency 5 10 20
"""
import asyncio
import argparse
import logging
import time
from datetime import datetime
from pathlib import Path
from typing import List
from .config import TestConfig
from .clients.asr_client import ASRWebSocketClient
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
# 配置日志
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
)
logger = logging.getLogger(__name__)
class ConcurrentBenchmark:
"""并发性能测试运行器"""
def __init__(self, config: TestConfig):
self.config = config
# 创建保存目录
self._setup_output_dirs()
def _setup_output_dirs(self):
"""创建输出目录结构"""
self.config.output_dir.mkdir(parents=True, exist_ok=True)
self.asr_output_dir = self.config.output_dir / "asr"
self.asr_output_dir.mkdir(exist_ok=True)
async def run_asr_benchmark(self) -> List[AggregatedMetrics]:
"""
运行 ASR 并发测试
Returns:
各并发级别的聚合指标列表
"""
logger.info("开始 ASR 并发性能测试...")
# 检查音频文件
if self.config.asr_audio_file is None:
raise ValueError("ASR 测试需要提供音频文件")
# 加载音频文件
audio_data, audio_duration = load_audio_file(
self.config.asr_audio_file,
self.config.asr_sample_rate,
)
audio_duration_ms = audio_duration * 1000
logger.info(f"音频文件已加载: {self.config.asr_audio_file}")
logger.info(f" - 时长: {audio_duration:.2f} 秒")
logger.info(f" - 大小: {len(audio_data) / 1024:.1f} KB")
results = []
for level in self.config.concurrency_levels:
logger.info(f"\n测试并发级别: {level}")
# 预热
logger.info(f" 预热中 ({self.config.warmup_requests} 次请求)...")
await self._run_asr_concurrent(
audio_data, audio_duration_ms, self.config.warmup_requests, level,
save_results=False
)
# 正式测试
logger.info(f" 正式测试中...")
start_time = time.perf_counter()
metrics_list = await self._run_asr_concurrent(
audio_data, audio_duration_ms, level, level,
save_results=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_asr_concurrent(
self,
audio_data: bytes,
audio_duration_ms: float,
num_requests: int,
concurrency_level: int,
save_results: bool = False,
) -> List[ASRMetrics]:
"""运行并发 ASR 请求"""
tasks = []
for _ in range(num_requests):
client = ASRWebSocketClient(
ws_url=self.config.asr_ws_url,
audio_data=audio_data,
audio_duration_ms=audio_duration_ms,
sample_rate=self.config.asr_sample_rate,
chunk_size=self.config.asr_chunk_size,
timeout=self.config.timeout_seconds,
save_result_dir=self.asr_output_dir if save_results else None,
)
tasks.append(client.run_test())
results = await asyncio.gather(*tasks, return_exceptions=True)
# 处理结果
metrics_list = []
for result in results:
if isinstance(result, ASRMetrics):
result.concurrency_level = concurrency_level
metrics_list.append(result)
else:
# 异常情况
metrics = ASRMetrics(
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],
) -> Path:
"""
生成测试报告
Args:
asr_results: ASR 测试结果
Returns:
报告文件路径
"""
output_dir = self.config.output_dir
output_dir.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
# 配置信息
config_info = {
"host": self.config.host,
"port": self.config.port,
"concurrency_levels": self.config.concurrency_levels,
}
# 生成 Markdown 报告
report_path = output_dir / f"{self.config.report_name}_{timestamp}.md"
reporter = MarkdownReporter()
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, output_dir, timestamp
)
for chart_file in chart_files:
logger.info(f"图表已生成: {chart_file}")
return report_path
def parse_args():
"""解析命令行参数"""
parser = argparse.ArgumentParser(
description="Qwen3-ASR 并发性能测试脚本",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
示例:
python -m scripts.benchmark.run --audio-file test.wav
# 自定义并发级别
python -m scripts.benchmark.run --audio-file test.wav --concurrency 5 10 20 50
""",
)
parser.add_argument(
"--host",
default="localhost",
help="服务器主机名 (默认: localhost)",
)
parser.add_argument(
"--port",
type=int,
default=8000,
help="服务器端口 (默认: 8000)",
)
parser.add_argument(
"--audio-file",
type=Path,
help="ASR 测试用音频文件路径(必需)",
)
parser.add_argument(
"--concurrency",
nargs="+",
type=int,
default=[5, 10, 20, 50],
help="并发级别列表 (默认: 5 10 20 50)",
)
parser.add_argument(
"--output",
type=Path,
default=Path("./benchmark_results"),
help="报告输出目录 (默认: ./benchmark_results)",
)
parser.add_argument(
"--timeout",
type=float,
default=120.0,
help="请求超时时间 (秒, 默认: 120)",
)
return parser.parse_args()
async def main():
"""主函数"""
args = parse_args()
# 创建配置
config = TestConfig(
host=args.host,
port=args.port,
concurrency_levels=args.concurrency,
asr_audio_file=args.audio_file,
output_dir=args.output,
timeout_seconds=args.timeout,
)
# 验证配置
try:
config.validate()
except ValueError as e:
logger.error(f"配置错误: {e}")
return
# 运行测试
benchmark = ConcurrentBenchmark(config)
asr_results = await benchmark.run_asr_benchmark()
# 生成报告
report_path = benchmark.generate_report(asr_results)
logger.info(f"\n测试完成! 报告已保存到: {report_path}")
if __name__ == "__main__":
asyncio.run(main())