| |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| from pathlib import Path |
|
|
| import matplotlib.pyplot as plt |
| import pandas as pd |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description="Compare experiment metrics") |
| parser.add_argument("--input", type=Path, default=Path("artifacts/experiments/urfd")) |
| args = parser.parse_args() |
| rows = [] |
| for path in sorted(args.input.glob("*/metrics.json")): |
| with path.open(encoding="utf-8") as file: |
| metrics = json.load(file) |
| rows.append( |
| { |
| "model": metrics["model"], |
| "accuracy": metrics["accuracy"], |
| "precision": metrics["precision"], |
| "recall": metrics["recall"], |
| "specificity": metrics["specificity"], |
| "f1": metrics["f1"], |
| "roc_auc": metrics["roc_auc"], |
| "training_seconds": metrics["training_seconds"], |
| } |
| ) |
| if not rows: |
| raise SystemExit(f"No metrics.json files found under {args.input}") |
| results = pd.DataFrame(rows).sort_values("f1", ascending=False) |
| results.to_csv(args.input / "results.csv", index=False) |
| results.set_index("model")[["accuracy", "precision", "recall", "specificity", "f1"]].plot( |
| kind="bar", figsize=(9, 5), ylim=(0, 1.05) |
| ) |
| plt.ylabel("Score") |
| plt.xlabel("Model") |
| plt.title("So sanh ket qua tren tap kiem tra") |
| plt.xticks(rotation=0) |
| plt.legend(loc="lower right", ncol=2) |
| plt.tight_layout() |
| plt.savefig(args.input / "model_comparison.png", dpi=180) |
| plt.close() |
| print(results.to_string(index=False, float_format=lambda value: f"{value:.4f}")) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|
|
|