Download scripts/run_repeated_experiments.py from minhy112/FallKLTN: direct link, hf CLI and curl.
- Browser
- Download file 4.57 kB
-
https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/run_repeated_experiments.py
- Command line
-
hf download hf://minhy112/FallKLTN/scripts/run_repeated_experiments.py
-
curl -L -o run_repeated_experiments.py https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/run_repeated_experiments.py
4.57 kB
| #!/usr/bin/env python3 | |
| 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() | |