移除TTS相关代码

main
Bifang 2026-09-30 17:36:04 +08:00
parent 5b4ab416e2
commit 86c9f1d415
13 changed files with 118 additions and 1059 deletions

View File

@ -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: 处理速度慢于实时,可能出现延迟累积

View File

@ -2,7 +2,7 @@
"""
Qwen3-ASR 并发性能测试脚本
用于测试 ASR/TTS WebSocket 服务在不同并发级别下的性能表现。
用于测试 ASR WebSocket 服务在不同并发级别下的性能表现。
"""
__version__ = "1.0.0"

View File

@ -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"]

View File

@ -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}")

View File

@ -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,22 +41,13 @@ 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():

View File

@ -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",

View File

@ -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

View File

@ -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])}")

View File

@ -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.grid(True, alpha=0.3)
if levels:
ax.set_xlabel("并发数", fontsize=12)
ax.set_ylabel(ylabel, fontsize=12)
ax.set_title(title, fontsize=14, fontweight="bold")
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 对比图"""
_fig, ax = plt.subplots(figsize=(10, 6))
if asr_results:
levels = [r.concurrency_level for r in asr_results]
avg_values = [r.rtf_avg for r in asr_results]
p95_values = [r.rtf_p95 for r in asr_results]
ax.plot(levels, avg_values, 'o-', color=self.colors["asr"],
label='ASR RTF (Avg)', linewidth=2, markersize=8)
ax.plot(levels, p95_values, 's--', color=self.colors["asr"],
label='ASR RTF (P95)', linewidth=1.5, markersize=6, alpha=0.7)
if tts_results:
levels = [r.concurrency_level for r in tts_results]
avg_values = [r.rtf_avg for r in tts_results]
p95_values = [r.rtf_p95 for r in tts_results]
ax.plot(levels, avg_values, 'o-', color=self.colors["tts"],
label='TTS RTF (Avg)', linewidth=2, markersize=8)
ax.plot(levels, p95_values, 's--', color=self.colors["tts"],
label='TTS RTF (P95)', linewidth=1.5, markersize=6, alpha=0.7)
# 添加 RTF=1.0 参考线
all_levels = set()
if asr_results:
all_levels.update(r.concurrency_level for r in asr_results)
if tts_results:
all_levels.update(r.concurrency_level for r in tts_results)
if all_levels:
ax.axhline(y=1.0, color='red', linestyle=':', linewidth=1.5,
label='RTF = 1.0 (实时)')
ax.set_xlabel('并发数', fontsize=12)
ax.set_ylabel('RTF', fontsize=12)
ax.set_title('RTF vs 并发数', fontsize=14, fontweight='bold')
ax.legend(loc='best')
ax.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:
"""生成吞吐量柱状图"""
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))
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()

View File

@ -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("")

View File

@ -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,18 +265,10 @@ 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()
# 生成报告
if asr_results or tts_results:
report_path = benchmark.generate_report(asr_results, tts_results)
report_path = benchmark.generate_report(asr_results)
logger.info(f"\n测试完成! 报告已保存到: {report_path}")

View File

@ -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"]

View File

@ -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()