-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain.py
More file actions
101 lines (77 loc) · 3.18 KB
/
Copy pathtrain.py
File metadata and controls
101 lines (77 loc) · 3.18 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
import argparse
import yaml
import datetime
import os
import glob
import tqdm
import utils
from dataset import dataloaders
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.")
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)
start = datetime.datetime.now()
configs['start'] = start
configs['version'] = version
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)
train_df = pd.read_csv(configs['train_df'])
val_df = pd.read_csv(configs['val_df'])
train_loader = dataloaders.get_train_dataloader(train_df, configs)
val_loader = dataloaders.get_validation_dataloader(val_df, configs)
model = ResNet18Classifier(pretrained=False)
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("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))
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))
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()