test/scripts/benchmark/reporters/chart_generator.py

122 lines
4.3 KiB
Python

# -*- coding: utf-8 -*-
"""Generate ASR benchmark charts."""
from pathlib import Path
from typing import List
import matplotlib
import matplotlib.pyplot as plt
import numpy as np
from ..metrics.models import AggregatedMetrics
# Keep generated labels readable on common Windows and Linux installations.
matplotlib.rcParams["font.sans-serif"] = ["Arial Unicode MS", "SimHei", "DejaVu Sans"]
matplotlib.rcParams["axes.unicode_minus"] = False
class ChartGenerator:
"""Create the latency, RTF, throughput, and duration ASR charts."""
_COLOR = "#4CAF50"
def generate_all_charts(
self,
results: List[AggregatedMetrics],
output_dir: Path,
timestamp: str,
) -> List[Path]:
output_dir.mkdir(parents=True, exist_ok=True)
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
def _save_line_chart(
self,
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]
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.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)
plt.tight_layout()
plt.savefig(output_path, dpi=150)
plt.close()
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))
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()