122 lines
4.3 KiB
Python
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()
|