File size: 4,567 Bytes
9313a90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
#!/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()