-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_cobre.py
More file actions
85 lines (70 loc) · 2.79 KB
/
Copy pathtrain_cobre.py
File metadata and controls
85 lines (70 loc) · 2.79 KB
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
"""Entry point for training LPAD on the COBRE dataset.
Example:
python scripts/train_cobre.py \
--data-root /path/to/COBRE \
--output-dir runs/lpad_cobre_seed42 \
--seed 42
"""
from __future__ import annotations
import argparse
import logging
import sys
from pathlib import Path
# Allow `python scripts/train_cobre.py` from repo root.
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from lpad import Config, Trainer # noqa: E402
def parse_args() -> Config:
parser = argparse.ArgumentParser(description="Train LPAD on COBRE.")
parser.add_argument("--data-root", type=Path, default=Config.data_root)
parser.add_argument("--output-dir", type=Path, default=Config.output_dir)
parser.add_argument("--epochs", type=int, default=Config.num_epochs)
parser.add_argument("--folds", type=int, default=Config.num_folds)
parser.add_argument("--batch-size", type=int, default=Config.batch_size)
parser.add_argument("--lr", type=float, default=Config.lr)
parser.add_argument("--seed", type=int, default=Config.seed)
parser.add_argument("--lambda-ortho", type=float, default=Config.lambda_ortho)
parser.add_argument("--lambda-supcon", type=float, default=Config.lambda_supcon)
parser.add_argument("--no-mixup", action="store_true")
parser.add_argument("--no-checkpoints", action="store_true")
parser.add_argument("--device", type=str, default=Config.device)
args = parser.parse_args()
return Config(
data_root=args.data_root,
output_dir=args.output_dir,
num_epochs=args.epochs,
num_folds=args.folds,
batch_size=args.batch_size,
lr=args.lr,
seed=args.seed,
lambda_ortho=args.lambda_ortho,
lambda_supcon=args.lambda_supcon,
mixup=not args.no_mixup,
save_checkpoints=not args.no_checkpoints,
device=args.device,
)
def main() -> None:
cfg = parse_args()
cfg.output_dir.mkdir(parents=True, exist_ok=True)
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(message)s",
handlers=[
logging.FileHandler(cfg.output_dir / "train.log"),
logging.StreamHandler(),
],
)
trainer = Trainer(cfg)
summary = trainer.run_cv()
# Pretty-print the headline result so it ends up at the bottom of stdout.
pearson = summary.mean_std("pearson")
mae = summary.mean_std("mae")
print("\n=== 5-fold cross-validation results ===")
print(f"{'Task':<6}{'Pearson':>20}{'MAE':>20}")
for task in cfg.task_names:
p_mean, p_std = pearson[task]
m_mean, m_std = mae[task]
print(f"{task:<6}{p_mean:>10.3f} ± {p_std:<6.3f}{m_mean:>10.3f} ± {m_std:<6.3f}")
if __name__ == "__main__":
main()