smartwatch-lm-0.2 / benchmark /benchmark_charts.py
prathamkode's picture
Make repo self-contained: rewrite docs, single-model benchmark, remove external references
dbb5d78 verified
Raw History Blame
3.84 kB
"""
Generate quality charts from a single-model benchmark report.
Usage:
python benchmark/benchmark_charts.py --report benchmark/report.json
python benchmark/benchmark_charts.py --report benchmark/report.json --output-dir benchmark/charts
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import matplotlib.pyplot as plt
import matplotlib.ticker as mticker
import numpy as np
BENCHMARK_DIR = Path(__file__).resolve().parent
DEFAULT_OUTPUT_DIR = BENCHMARK_DIR / "charts"
COLOR = "#55A868"
METRIC_LABELS = [
("intent_accuracy", "Intent accuracy"),
("intent_parse_rate", "Intent parse rate"),
("clean_output_rate", "Clean output rate"),
("slot_presence", "Slot presence"),
]
def _pct_axis(ax) -> None:
ax.yaxis.set_major_formatter(mticker.PercentFormatter(xmax=1.0, decimals=0))
ax.set_ylim(0, 1.05)
def _model_from_report(data: dict) -> dict:
models = data.get("models") or []
if len(models) != 1:
raise ValueError("Report must contain exactly one model.")
return models[0]
def save_overall_metrics_chart(model: dict, output_path: Path) -> None:
labels = [label for _, label in METRIC_LABELS]
keys = [key for key, _ in METRIC_LABELS]
values = [model[key] for key in keys]
x = np.arange(len(labels))
width = 0.5
fig, ax = plt.subplots(figsize=(9, 5))
bars = ax.bar(x, values, width, label=model["name"], color=COLOR)
ax.set_title("Overall Quality Metrics")
ax.set_xticks(x)
ax.set_xticklabels(labels, rotation=15, ha="right")
ax.set_ylabel("Score")
_pct_axis(ax)
ax.legend(loc="lower right")
ax.grid(axis="y", alpha=0.3)
ax.bar_label(bars, fmt="%.0f%%", padding=2, labels=[f"{v * 100:.0f}%" for v in values])
fig.tight_layout()
fig.savefig(output_path, dpi=150)
plt.close(fig)
def save_per_intent_chart(model: dict, output_path: Path) -> None:
per_intent = model.get("per_intent") or {}
intents = sorted(per_intent)
accuracies = [per_intent[i].get("accuracy", 0.0) for i in intents]
y = np.arange(len(intents))
fig_h = max(6, len(intents) * 0.28)
fig, ax = plt.subplots(figsize=(10, fig_h))
ax.barh(y, accuracies, color=COLOR)
ax.set_title("Per-Intent Accuracy")
ax.set_yticks(y)
ax.set_yticklabels(intents, fontsize=8)
ax.set_xlabel("Accuracy")
ax.invert_yaxis()
ax.xaxis.set_major_formatter(mticker.PercentFormatter(xmax=1.0, decimals=0))
ax.set_xlim(0, 1.05)
ax.grid(axis="x", alpha=0.3)
fig.tight_layout()
fig.savefig(output_path, dpi=150)
plt.close(fig)
def save_charts_from_report(report_path: Path, output_dir: Path) -> list[Path]:
data = json.loads(report_path.read_text(encoding="utf-8"))
model = _model_from_report(data)
output_dir.mkdir(parents=True, exist_ok=True)
charts = [
output_dir / "overall_metrics.png",
output_dir / "per_intent_accuracy.png",
]
save_overall_metrics_chart(model, charts[0])
save_per_intent_chart(model, charts[1])
return charts
def main() -> None:
parser = argparse.ArgumentParser(description="Generate benchmark charts for Smartwatch LM v0.2")
parser.add_argument("--report", type=Path, required=True, help="JSON report")
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
args = parser.parse_args()
if not args.report.is_file():
raise SystemExit(f"Report not found: {args.report}")
charts = save_charts_from_report(args.report, args.output_dir)
print(f"Wrote {len(charts)} charts -> {args.output_dir}")
for path in charts:
print(f" {path.name}")
if __name__ == "__main__":
main()