# -*- 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()