Make repo self-contained: rewrite docs, single-model benchmark, remove external references
dbb5d78 verified Download benchmark/benchmark_charts.py from prathamkode/smartwatch-lm-0.2: direct link, hf CLI and curl.
- Browser
- Download file 3.84 kB
-
https://huggingface.co/prathamkode/smartwatch-lm-0.2/resolve/dbb5d78326ae2eace84b67a293deaf2dcc815dab/benchmark/benchmark_charts.py
- Command line
-
hf download hf://prathamkode/smartwatch-lm-0.2@dbb5d78326ae2eace84b67a293deaf2dcc815dab/benchmark/benchmark_charts.py
-
curl -L -o benchmark_charts.py https://huggingface.co/prathamkode/smartwatch-lm-0.2/resolve/dbb5d78326ae2eace84b67a293deaf2dcc815dab/benchmark/benchmark_charts.py
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() | |