From 86c9f1d41582240c70cb8fd4ebb2e0e9dcfdeeb6 Mon Sep 17 00:00:00 2001 From: Bifang <915779419@qq.com> Date: Wed, 30 Sep 2026 17:36:04 +0800 Subject: [PATCH] =?UTF-8?q?=E7=A7=BB=E9=99=A4TTS=E7=9B=B8=E5=85=B3?= =?UTF-8?q?=E4=BB=A3=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/benchmark/README.md | 49 +-- scripts/benchmark/__init__.py | 2 +- scripts/benchmark/clients/__init__.py | 3 +- scripts/benchmark/clients/tts_client.py | 251 --------------- scripts/benchmark/config.py | 29 +- scripts/benchmark/metrics/__init__.py | 3 +- scripts/benchmark/metrics/models.py | 54 ---- scripts/benchmark/metrics/statistics.py | 97 +----- .../benchmark/reporters/chart_generator.py | 295 +++++------------- .../benchmark/reporters/markdown_reporter.py | 61 +--- scripts/benchmark/run.py | 160 +--------- scripts/benchmark/utils/__init__.py | 3 +- scripts/benchmark/utils/text_generator.py | 170 ---------- 13 files changed, 118 insertions(+), 1059 deletions(-) delete mode 100644 scripts/benchmark/clients/tts_client.py delete mode 100644 scripts/benchmark/utils/text_generator.py diff --git a/scripts/benchmark/README.md b/scripts/benchmark/README.md index 2a5548e..dcc24d3 100644 --- a/scripts/benchmark/README.md +++ b/scripts/benchmark/README.md @@ -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: 处理速度慢于实时,可能出现延迟累积 diff --git a/scripts/benchmark/__init__.py b/scripts/benchmark/__init__.py index 35afc3c..b016eb4 100644 --- a/scripts/benchmark/__init__.py +++ b/scripts/benchmark/__init__.py @@ -2,7 +2,7 @@ """ Qwen3-ASR 并发性能测试脚本 -用于测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。 +用于测试 ASR WebSocket 服务在不同并发级别下的性能表现。 """ __version__ = "1.0.0" diff --git a/scripts/benchmark/clients/__init__.py b/scripts/benchmark/clients/__init__.py index 46938a9..bb13602 100644 --- a/scripts/benchmark/clients/__init__.py +++ b/scripts/benchmark/clients/__init__.py @@ -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"] diff --git a/scripts/benchmark/clients/tts_client.py b/scripts/benchmark/clients/tts_client.py deleted file mode 100644 index 0ef6075..0000000 --- a/scripts/benchmark/clients/tts_client.py +++ /dev/null @@ -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}") diff --git a/scripts/benchmark/config.py b/scripts/benchmark/config.py index bb1341a..416b8fc 100644 --- a/scripts/benchmark/config.py +++ b/scripts/benchmark/config.py @@ -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("至少需要一个并发级别") diff --git a/scripts/benchmark/metrics/__init__.py b/scripts/benchmark/metrics/__init__.py index c3a8551..9b408fa 100644 --- a/scripts/benchmark/metrics/__init__.py +++ b/scripts/benchmark/metrics/__init__.py @@ -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", diff --git a/scripts/benchmark/metrics/models.py b/scripts/benchmark/metrics/models.py index 0aa7e62..1382e95 100644 --- a/scripts/benchmark/metrics/models.py +++ b/scripts/benchmark/metrics/models.py @@ -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 diff --git a/scripts/benchmark/metrics/statistics.py b/scripts/benchmark/metrics/statistics.py index 69e8da5..b2eb2e5 100644 --- a/scripts/benchmark/metrics/statistics.py +++ b/scripts/benchmark/metrics/statistics.py @@ -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])}") diff --git a/scripts/benchmark/reporters/chart_generator.py b/scripts/benchmark/reporters/chart_generator.py index 57b367d..782c15d 100644 --- a/scripts/benchmark/reporters/chart_generator.py +++ b/scripts/benchmark/reporters/chart_generator.py @@ -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() diff --git a/scripts/benchmark/reporters/markdown_reporter.py b/scripts/benchmark/reporters/markdown_reporter.py index ab93b32..3fa9c19 100644 --- a/scripts/benchmark/reporters/markdown_reporter.py +++ b/scripts/benchmark/reporters/markdown_reporter.py @@ -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("") diff --git a/scripts/benchmark/run.py b/scripts/benchmark/run.py index 38d0065..bdf6eed 100644 --- a/scripts/benchmark/run.py +++ b/scripts/benchmark/run.py @@ -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__": diff --git a/scripts/benchmark/utils/__init__.py b/scripts/benchmark/utils/__init__.py index 68c179e..d22898e 100644 --- a/scripts/benchmark/utils/__init__.py +++ b/scripts/benchmark/utils/__init__.py @@ -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"] diff --git a/scripts/benchmark/utils/text_generator.py b/scripts/benchmark/utils/text_generator.py deleted file mode 100644 index dececb4..0000000 --- a/scripts/benchmark/utils/text_generator.py +++ /dev/null @@ -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()