FallKLTN / scripts /run_repeated_experiments.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw
History Blame Contribute Delete
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()