-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
52 lines (47 loc) · 1.81 KB
/
Copy pathmain.py
File metadata and controls
52 lines (47 loc) · 1.81 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
import os
from datetime import datetime
from zipfile import ZipFile
from dotenv import load_dotenv
from Config import Config
from Evaluation.training_validation import training_validation
from Evaluation.utilities import check_cuda_availability
from Models.model_adjustment import adjust
from training_iterables import spectrogram_training_iterable
if __name__ == "__main__":
load_dotenv()
if not Config.vowels_path.exists():
print("Downloading data...")
zip_path = Config.data_path / "Vowels.zip"
os.system(
f"gdown https://drive.google.com/uc?id={Config.google_drive_file_id} -O {zip_path}"
)
with ZipFile(zip_path) as zip_ref:
zip_ref.extractall(Config.data_path)
os.remove(zip_path)
print("Files downloaded")
if Config.device is None:
Config.device = check_cuda_availability()
project_name = datetime.now().strftime('%Y%m%d%H%M')
for (
model_creation_function,
vowels,
(window_arguments, augmentation),
) in spectrogram_training_iterable:
many_channels = vowels[0] == "all"
model_creation_function = adjust(
model_creation_function, many_channels, *window_arguments
)
training_validation(
device=Config.device,
vowels=vowels,
batch_size=Config.batch_size,
num_splits=Config.num_splits,
early_stopping_patience=Config.early_stopping_patience,
criterion=Config.criterion,
model_creator=model_creation_function,
learning_rate=Config.learning_rate,
learning_rate_scheduler_creator=Config.learning_rate_scheduler_creator,
random_state=Config.random_state,
augmentation=augmentation,
project_name=project_name,
)