-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun_train.py
More file actions
85 lines (68 loc) · 3.3 KB
/
Copy pathrun_train.py
File metadata and controls
85 lines (68 loc) · 3.3 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
"""
run_train.py - Orchestrates and launches CABT handwriting decoder training.
Resolves directory paths, handles checkpoint auto-resuming, and invokes training cycles.
"""
import os
import sys
import pickle
from datetime import datetime
# Add local path to PYTHONPATH
repo_dir = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, repo_dir)
from cabt import get_default_decoder_args, CABT
# ---------------------------------------------------------------------------
# Training Configurations
# ---------------------------------------------------------------------------
root_dir = os.path.expanduser("~") + "/handwritingBCIData/"
data_dirs = [
"t5.2019.05.08", "t5.2019.11.25", "t5.2019.12.09", "t5.2019.12.11",
"t5.2019.12.18", "t5.2019.12.20", "t5.2020.01.06", "t5.2020.01.08",
"t5.2020.01.13", "t5.2020.01.15",
]
cv_part = "HeldOutTrials" # Default evaluation partition (HeldOutTrials or HeldOutBlocks)
output_dir = os.path.join(root_dir, "RNNTrainingSteps/Step4_RNNTraining/", cv_part)
os.makedirs(output_dir, exist_ok=True)
# Instantiate default model arguments
args = get_default_decoder_args()
# Configure data files for multi-day calibration
for day_idx, directory in enumerate(data_dirs):
args[f"sentencesFile_{day_idx}"] = os.path.join(root_dir, "Datasets/", directory, "sentences.mat")
args[f"singleLettersFile_{day_idx}"] = os.path.join(root_dir, "Datasets/", directory, "singleLetters.mat")
args[f"labelsFile_{day_idx}"] = os.path.join(
root_dir, "RNNTrainingSteps/Step2_HMMLabels/", cv_part, f"{directory}_timeSeriesLabels.mat"
)
args[f"syntheticDatasetDir_{day_idx}"] = os.path.join(
root_dir, "RNNTrainingSteps/Step3_SyntheticSentences/", cv_part, f"{directory}_syntheticSentences/"
)
args[f"cvPartitionFile_{day_idx}"] = os.path.join(
root_dir, "RNNTrainingSteps/", f"trainTestPartitions_{cv_part}.mat"
)
args[f"sessionName_{day_idx}"] = directory
args["outputDir"] = output_dir
args["dayProbability"] = "[0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1, 0.1]"
args["dayToLayerMap"] = "[0, 1, 2, 3, 4, 5, 6, 7, 8, 9]"
args["mode"] = "train"
# ---------------------------------------------------------------------------
# Auto-Resume Checkpoint Setup
# ---------------------------------------------------------------------------
checkpoint_index = os.path.join(output_dir, "checkpoint")
if os.path.isfile(checkpoint_index):
print(f"[run_train] Checkpoint found in {output_dir}. Auto-resuming training.")
args["loadDir"] = output_dir
else:
print(f"[run_train] No checkpoints detected. Starting training fresh.")
args["loadDir"] = "None"
# Serialize config dictionary
args_file = os.path.join(output_dir, "args.p")
pickle.dump(args, open(args_file, "wb"))
print(f"[run_train] Model arguments saved -> {args_file}")
# ---------------------------------------------------------------------------
# Execute Training
# ---------------------------------------------------------------------------
print(f"[run_train] Instantiating CABT model ({args['nBatchesToTrain']:,} batches total).")
print(f" Checkpointing frequency: {args['batchesPerModelSave']:,} batches.")
print(f" Max checkpoints to keep: {args['nCheckToKeep']}.")
print()
model = CABT(args=args)
model.train()
print("\n[run_train] Training process completed successfully.")