| |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import subprocess |
| import sys |
| from pathlib import Path |
|
|
| import matplotlib.pyplot as plt |
| import pandas as pd |
|
|
| METRICS = ["accuracy", "precision", "recall", "specificity", "f1", "roc_auc"] |
|
|
|
|
| def run(command: list[str]) -> None: |
| print("\n$", " ".join(command), flush=True) |
| subprocess.run(command, check=True) |
|
|
|
|
| def collect_results(output: Path, seeds: list[int]) -> pd.DataFrame: |
| rows = [] |
| for seed in seeds: |
| run_dir = output / f"seed_{seed}" |
| for metrics_path in sorted(run_dir.glob("*/metrics.json")): |
| with metrics_path.open(encoding="utf-8") as file: |
| metrics = json.load(file) |
| row = {"seed": seed, "model": metrics["model"]} |
| row.update({metric: metrics[metric] for metric in METRICS}) |
| row["training_seconds"] = metrics["training_seconds"] |
| row["device"] = metrics.get("device", "cpu") |
| rows.append(row) |
| return pd.DataFrame(rows) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser( |
| description="Repeat real-data experiments over multiple split/training seeds" |
| ) |
| parser.add_argument("--dataset", default="data/processed/urfd_pose.npz") |
| parser.add_argument("--config", default="configs/default.yaml") |
| parser.add_argument( |
| "--output", type=Path, default=Path("artifacts/experiments/urfd_repeated") |
| ) |
| parser.add_argument("--seeds", type=int, nargs="+", default=[13, 21, 42, 84, 123]) |
| parser.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") |
| parser.add_argument( |
| "--skip-existing", action="store_true", help="Reuse a seed if all four metrics files exist" |
| ) |
| args = parser.parse_args() |
| if not Path(args.dataset).exists(): |
| raise SystemExit( |
| f"Real dataset not found: {args.dataset}. Run scripts/download_urfd.py and " |
| "scripts/prepare_dataset.py first." |
| ) |
| args.output.mkdir(parents=True, exist_ok=True) |
| python = sys.executable |
|
|
| for seed in args.seeds: |
| run_dir = args.output / f"seed_{seed}" |
| expected = [ |
| run_dir / model / "metrics.json" |
| for model in ("logistic_regression", "random_forest", "pose_gru", "pose_tcn") |
| ] |
| if args.skip_existing and all(path.exists() for path in expected): |
| print(f"Reusing complete run for seed {seed}") |
| continue |
| common = [ |
| "--dataset", args.dataset, |
| "--config", args.config, |
| "--output", str(run_dir), |
| "--seed", str(seed), |
| ] |
| run([python, "scripts/train_baselines.py", *common]) |
| run([python, "scripts/train_gru.py", *common, "--device", args.device]) |
| run([python, "scripts/train_tcn.py", *common, "--device", args.device]) |
| run([python, "scripts/compare_models.py", "--input", str(run_dir)]) |
|
|
| per_run = collect_results(args.output, args.seeds) |
| expected_rows = len(args.seeds) * 4 |
| if len(per_run) != expected_rows: |
| raise RuntimeError(f"Expected {expected_rows} result rows, found {len(per_run)}") |
| per_run.to_csv(args.output / "per_run_results.csv", index=False) |
|
|
| aggregate = per_run.groupby("model")[METRICS + ["training_seconds"]].agg( |
| ["mean", "std", "min", "max"] |
| ) |
| aggregate.columns = [f"{metric}_{stat}" for metric, stat in aggregate.columns] |
| aggregate = aggregate.reset_index().sort_values("f1_mean", ascending=False) |
| aggregate.to_csv(args.output / "aggregate_results.csv", index=False) |
|
|
| plot_data = per_run.pivot(index="seed", columns="model", values="f1") |
| means = plot_data.mean().sort_values(ascending=False) |
| stds = plot_data.std().reindex(means.index) |
| plt.figure(figsize=(8, 5)) |
| plt.bar(means.index, means.values, yerr=stds.values, capsize=5) |
| plt.ylim(0, 1.05) |
| plt.ylabel("F1 (mean ± standard deviation)") |
| plt.xlabel("Model") |
| plt.title(f"URFD repeated experiment ({len(args.seeds)} seeds)") |
| plt.xticks(rotation=10) |
| plt.tight_layout() |
| plt.savefig(args.output / "aggregate_f1.png", dpi=180) |
| plt.close() |
|
|
| print("\nPer-run results:") |
| print(per_run.to_string(index=False, float_format=lambda value: f"{value:.4f}")) |
| print("\nAggregate results:") |
| display_columns = ["model"] + [f"{metric}_{stat}" for metric in METRICS for stat in ("mean", "std")] |
| print(aggregate[display_columns].to_string(index=False, float_format=lambda value: f"{value:.4f}")) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|