File size: 1,476 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 | #!/usr/bin/env python3
from __future__ import annotations
import argparse
import subprocess
import sys
from pathlib import Path
def run(command: list[str]) -> None:
print("\n$", " ".join(command), flush=True)
subprocess.run(command, check=True)
def main() -> None:
parser = argparse.ArgumentParser(description="Run the complete experiment")
parser.add_argument("--dataset", default="data/processed/urfd_pose.npz")
parser.add_argument("--config", default="configs/default.yaml")
parser.add_argument("--output", default="artifacts/experiments/urfd")
parser.add_argument("--seed", type=int)
parser.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto")
args = parser.parse_args()
python = sys.executable
dataset = args.dataset
output = args.output
if not Path(dataset).exists():
run([python, "scripts/download_urfd.py"])
run([python, "scripts/prepare_dataset.py", "--config", args.config, "--output", dataset])
common = ["--dataset", dataset, "--config", args.config, "--output", output]
if args.seed is not None:
common.extend(["--seed", str(args.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", output])
if __name__ == "__main__":
main()
|