-
Notifications
You must be signed in to change notification settings - Fork 13
Expand file tree
/
Copy pathtrain.py
More file actions
56 lines (44 loc) · 1.76 KB
/
Copy pathtrain.py
File metadata and controls
56 lines (44 loc) · 1.76 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
from validate import validate
from data import create_dataloader
from trainer.trainer import Trainer
from options.train_options import TrainOptions
def get_val_opt():
val_opt = TrainOptions().parse(print_options=False)
val_opt.isTrain = False
val_opt.data_label = "val"
val_opt.real_list_path = "./datasets/val/0_real"
val_opt.fake_list_path = "./datasets/val/1_fake"
return val_opt
if __name__ == "__main__":
opt = TrainOptions().parse()
val_opt = get_val_opt()
model = Trainer(opt)
data_loader = create_dataloader(opt)
val_loader = create_dataloader(val_opt)
print("Length of data loader: %d" % (len(data_loader)))
print("Length of val loader: %d" % (len(val_loader)))
for epoch in range(opt.epoch):
model.train()
print("epoch: ", epoch + model.step_bias)
for i, (img, crops, label) in enumerate(data_loader):
model.total_steps += 1
model.set_input((img, crops, label))
model.forward()
loss = model.get_loss()
model.optimize_parameters()
if model.total_steps % opt.loss_freq == 0:
print(
"Train loss: {}\tstep: {}".format(
model.get_loss(), model.total_steps
)
)
if epoch % opt.save_epoch_freq == 0:
print("saving the model at the end of epoch %d" % (epoch + model.step_bias))
model.save_trainer("model_epoch_%s.pth" % (epoch + model.step_bias))
model.eval()
ap, fpr, fnr, acc = validate(model.model, val_loader, opt.gpu_ids)
print(
"(Val @ epoch {}) acc: {} ap: {} fpr: {} fnr: {}".format(
epoch + model.step_bias, acc, ap, fpr, fnr
)
)