-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain_fasd.py
More file actions
113 lines (86 loc) · 3.52 KB
/
Copy pathtrain_fasd.py
File metadata and controls
113 lines (86 loc) · 3.52 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
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
import argparse
import yaml
import datetime
import os
import glob
import tqdm
import utils
from dataset import casia_fasd
from models.scan import SCAN, Vigilant, SCANEncoder
from models.resnet import ResNet18Classifier
from vigilant import *
import torch
import torchvision
from torch.utils.tensorboard import SummaryWriter
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("-c", "--config", required=True, help="Config file path.")
parser.add_argument("-ch", "--checkpoint", required=False, help="Checkpoint file path.", default=None)
args = parser.parse_args()
with open(args.config, 'r') as stream:
configs = yaml.safe_load(stream)
root_dir = configs['log_dir']
if not os.path.isdir(root_dir):
os.makedirs(root_dir)
version = utils.get_latest_version(root_dir)
version_directory = root_dir + "version_" + str(version)
if not os.path.isdir(version_directory):
os.makedirs(version_directory)
weights_directory = version_directory + "/weights/"
if not os.path.isdir(weights_directory):
os.makedirs(weights_directory)
debug = False
start = datetime.datetime.now()
configs['start'] = start
configs['version'] = version
configs['debug'] = debug
with open(version_directory + '/configs.yml', 'w') as outfile:
yaml.dump(configs, outfile, default_flow_style=False)
# ========================= End of DevOps ==========================
# ========================= Start of ML ==========================
device = configs['device']
print("Using", device)
print("Version ", version)
print("Debug: ", debug)
train_df = pd.read_csv(configs['train_df'])
val_df = pd.read_csv(configs['val_df'])
train_loader = casia_fasd.get_train_dataloader(train_df, configs)
val_loader = casia_fasd.get_validation_dataloader(val_df, configs)
model = ResNet18Classifier(pretrained=False)
try:
model.load_state_dict(configs['checkpoint'])
except KeyError:
pass
model.to(device)
optim = torch.optim.Adam(model.parameters(), lr=configs['lr'])
criterion = torch.nn.CrossEntropyLoss()
writer = SummaryWriter(log_dir=version_directory)
if configs['print_model']:
print(model)
print("====================")
print("Starting training...")
training_avg_losses = []
val_avg_losses = []
for i in tqdm.trange(int(configs['max_epochs']), desc="Epoch"):
training_losses = train(model, device, optim, criterion, train_loader,
writer, int(i + 1), configs, debug)
train_avg_loss = np.mean(training_losses)
training_avg_losses.append(train_avg_loss)
writer.add_scalar('Loss (epoch)/Training average loss', train_avg_loss, i)
with torch.no_grad():
val_losses = validate(model, device, optim, criterion, val_loader,
writer, int(i + 1), configs, debug)
val_avg_loss = np.mean(val_losses)
val_avg_losses.append(val_avg_loss)
writer.add_scalar('Loss (epoch)/Validation average loss', val_avg_loss, i)
torch.save(model.state_dict(), weights_directory + "epoch_" + str(i) + ".pth")
plt.plot(training_avg_losses, label="Training average loss")
plt.plot(val_avg_losses, label="Validation average loss")
plt.title("Overall Loss curve")
plt.legend()
plt.grid()
plt.savefig(version_directory + "/losses.png")
writer.close()